use crate::tagger_data::{head_a_labels, head_b_labels, to_bio, to_epistemic, TaggerExample};
use crate::vocabulary::VocabularySpace;
use candle_core::{DType, Device, IndexOp, Tensor, D};
use candle_nn::{loss, ops::softmax, AdamW, Linear, Module, Optimizer, ParamsAdamW, VarBuilder, VarMap};
use crate::hrm::{HrmConfig, HrmTagger};
use candle_transformers::models::bert::Config as BertConfig;
use serde::Serialize;
use std::path::{Path, PathBuf};
use tokenizers::Tokenizer;
const IGNORE: i64 = -100;
#[derive(Debug, Clone)]
pub struct TrainConfig {
pub base_dir: PathBuf,
pub tokenizer: PathBuf,
pub epochs: usize,
pub lr: f64,
pub batch: usize,
pub max_len: usize,
pub lambda_b: f64,
pub seed: u64,
}
impl Default for TrainConfig {
fn default() -> Self {
TrainConfig {
base_dir: PathBuf::new(),
tokenizer: PathBuf::new(),
epochs: 8,
lr: 5e-4,
batch: 8,
max_len: 128,
lambda_b: 0.5,
seed: 0xC0FFEE,
}
}
}
pub type MultiHeadTagger = HrmTagger;
pub fn hrm_config_from(bert: &BertConfig) -> HrmConfig {
HrmConfig::bert_tiny(bert.vocab_size, bert.hidden_size, bert.max_position_embeddings, bert.type_vocab_size)
}
#[derive(Debug, Clone)]
pub struct Encoded {
pub ids: Vec<u32>,
pub attn: Vec<u32>,
pub labels_a: Vec<i64>,
pub labels_b: Vec<i64>,
}
pub fn encode(spec: &VocabularySpace, tok: &Tokenizer, ex: &TaggerExample, max_len: usize) -> Result<Encoded, String> {
let enc = tok.encode(ex.text.as_str(), true).map_err(|e| format!("tokenize: {e}"))?;
let offsets: Vec<(usize, usize)> = enc.get_offsets().to_vec();
let mut ids: Vec<u32> = enc.get_ids().to_vec();
let mut labels_a = to_bio(spec, &ex.spans, &offsets);
let mut labels_b = to_epistemic(&ex.spans, &offsets);
let n = ids.len().min(max_len);
ids.truncate(n);
labels_a.truncate(n);
labels_b.truncate(n);
let mut attn = vec![1u32; n];
while ids.len() < max_len {
ids.push(0);
attn.push(0);
labels_a.push(IGNORE);
labels_b.push(IGNORE);
}
Ok(Encoded { ids, attn, labels_a, labels_b })
}
fn masked_ce(logits: &Tensor, labels: &[i64], device: &Device) -> candle_core::Result<Option<Tensor>> {
let keep: Vec<u32> = labels.iter().enumerate().filter(|(_, l)| **l != IGNORE).map(|(i, _)| i as u32).collect();
if keep.is_empty() {
return Ok(None);
}
let (b, t, c) = logits.dims3()?;
let flat = logits.reshape((b * t, c))?;
let idx = Tensor::from_vec(keep.clone(), keep.len(), device)?;
let picked = flat.index_select(&idx, 0)?;
let tgt: Vec<u32> = keep.iter().map(|&i| labels[i as usize] as u32).collect();
let tgt = Tensor::from_vec(tgt, keep.len(), device)?;
Ok(Some(loss::cross_entropy(&picked, &tgt)?))
}
#[derive(Debug, Clone, Serialize)]
pub struct EpochReport {
pub epoch: usize,
pub loss_a: f64,
pub loss_b: f64,
pub total: f64,
}
#[derive(Debug, Clone, Serialize)]
pub struct TrainReport {
pub examples: usize,
pub labels_a: usize,
pub labels_b: usize,
pub epochs: Vec<EpochReport>,
pub train_acc_a: f64,
pub train_acc_b: f64,
pub dev_examples: usize,
pub dev_acc_a: f64,
pub dev_acc_b: f64,
pub dev_acc_a_spans: f64,
}
pub fn train(
spec: &VocabularySpace,
examples: &[TaggerExample],
cfg: &TrainConfig,
) -> Result<(VarMap, TrainReport, Vec<String>), String> {
if examples.is_empty() {
return Err("no training examples".into());
}
let device = Device::Cpu;
let bert_cfg: BertConfig = serde_json::from_slice(
&std::fs::read(cfg.base_dir.join("config.json")).map_err(|e| format!("base config.json: {e}"))?,
)
.map_err(|e| format!("parse base config: {e}"))?;
let tok = Tokenizer::from_file(&cfg.tokenizer).map_err(|e| format!("tokenizer: {e}"))?;
let labels_a = head_a_labels(spec);
let n_a = labels_a.len();
let n_b = head_b_labels().len();
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let hrm_cfg = hrm_config_from(&bert_cfg);
let model = HrmTagger::new(vb, &hrm_cfg, n_a, n_b).map_err(|e| format!("build model: {e}"))?;
{
let weights = cfg.base_dir.join("model.safetensors");
let pre = candle_core::safetensors::load(&weights, &device).map_err(|e| format!("load {}: {e}", weights.display()))?;
let mut vm = varmap.clone();
let mut loaded = 0usize;
for (name, t) in pre.iter().filter(|(n, _)| n.starts_with("bert.embeddings.")) {
if vm.set_one(name, t).is_ok() {
loaded += 1;
}
}
if loaded == 0 {
return Err("no pretrained embedding tensors matched (expected `bert.embeddings.*`)".into());
}
eprintln!("loaded {loaded} pretrained embedding tensors; HRM core + heads train from scratch");
}
let all: Vec<Encoded> = examples.iter().filter_map(|e| encode(spec, &tok, e, cfg.max_len).ok()).collect();
if all.len() < 5 {
return Err("too few examples to split train/dev".into());
}
let mut encoded: Vec<Encoded> = Vec::new();
let mut dev: Vec<Encoded> = Vec::new();
for (i, e) in all.into_iter().enumerate() {
if i % 5 == 4 {
dev.push(e);
} else {
encoded.push(e);
}
}
let mut opt = AdamW::new(varmap.all_vars(), ParamsAdamW { lr: cfg.lr, ..Default::default() })
.map_err(|e| format!("optimizer: {e}"))?;
let mut epochs_out = Vec::new();
let mut order: Vec<usize> = (0..encoded.len()).collect();
let mut seed = cfg.seed;
let mut next = move || {
seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
(seed >> 33) as usize
};
for epoch in 1..=cfg.epochs {
for i in (1..order.len()).rev() {
order.swap(i, next() % (i + 1));
}
let (mut sum_a, mut sum_b, mut steps) = (0.0f64, 0.0f64, 0usize);
for chunk in order.chunks(cfg.batch.max(1)) {
let bs = chunk.len();
let ids: Vec<u32> = chunk.iter().flat_map(|&i| encoded[i].ids.clone()).collect();
let attn: Vec<u32> = chunk.iter().flat_map(|&i| encoded[i].attn.clone()).collect();
let la: Vec<i64> = chunk.iter().flat_map(|&i| encoded[i].labels_a.clone()).collect();
let lb: Vec<i64> = chunk.iter().flat_map(|&i| encoded[i].labels_b.clone()).collect();
let ids_t = Tensor::from_vec(ids, (bs, cfg.max_len), &device).map_err(|e| e.to_string())?;
let attn_t = Tensor::from_vec(attn, (bs, cfg.max_len), &device).map_err(|e| e.to_string())?;
let (log_a, log_b) = model.forward(&ids_t, &attn_t, true).map_err(|e| format!("forward: {e}"))?;
let ce_a = masked_ce(&log_a, &la, &device).map_err(|e| e.to_string())?;
let ce_b = masked_ce(&log_b, &lb, &device).map_err(|e| e.to_string())?;
let (Some(ce_a), Some(ce_b)) = (ce_a, ce_b) else { continue };
let total = (&ce_a + (ce_b.affine(cfg.lambda_b, 0.0).map_err(|e| e.to_string())?)).map_err(|e| e.to_string())?;
opt.backward_step(&total).map_err(|e| format!("backward: {e}"))?;
sum_a += ce_a.to_scalar::<f32>().map_err(|e| e.to_string())? as f64;
sum_b += ce_b.to_scalar::<f32>().map_err(|e| e.to_string())? as f64;
steps += 1;
}
let d = steps.max(1) as f64;
epochs_out.push(EpochReport {
epoch,
loss_a: round4(sum_a / d),
loss_b: round4(sum_b / d),
total: round4((sum_a + cfg.lambda_b * sum_b) / d),
});
}
let (acc_a, acc_b, _) = accuracy(&model, &encoded, cfg, &device).map_err(|e| e.to_string())?;
let (dev_a, dev_b, dev_span) = accuracy(&model, &dev, cfg, &device).map_err(|e| e.to_string())?;
let report = TrainReport {
examples: encoded.len(),
labels_a: n_a,
labels_b: n_b,
epochs: epochs_out,
train_acc_a: round4(acc_a),
train_acc_b: round4(acc_b),
dev_examples: dev.len(),
dev_acc_a: round4(dev_a),
dev_acc_b: round4(dev_b),
dev_acc_a_spans: round4(dev_span),
};
Ok((varmap, report, labels_a))
}
fn accuracy(model: &MultiHeadTagger, encoded: &[Encoded], cfg: &TrainConfig, device: &Device) -> candle_core::Result<(f64, f64, f64)> {
let (mut ok_a, mut ok_b, mut n_a, mut n_b) = (0usize, 0usize, 0usize, 0usize);
let (mut ok_span, mut n_span) = (0usize, 0usize);
for e in encoded {
let ids = Tensor::from_vec(e.ids.clone(), (1, cfg.max_len), device)?;
let attn = Tensor::from_vec(e.attn.clone(), (1, cfg.max_len), device)?;
let (la, lb) = model.forward(&ids, &attn, false)?;
let pa = softmax(&la.i(0)?, D::Minus1)?.argmax(D::Minus1)?.to_vec1::<u32>()?;
let pb = softmax(&lb.i(0)?, D::Minus1)?.argmax(D::Minus1)?.to_vec1::<u32>()?;
for (i, &t) in e.labels_a.iter().enumerate() {
if t != IGNORE {
n_a += 1;
if pa[i] as i64 == t {
ok_a += 1;
}
if t != 0 {
n_span += 1;
if pa[i] as i64 == t {
ok_span += 1;
}
}
}
}
for (i, &t) in e.labels_b.iter().enumerate() {
if t != IGNORE {
n_b += 1;
if pb[i] as i64 == t {
ok_b += 1;
}
}
}
}
Ok((
ok_a as f64 / n_a.max(1) as f64,
ok_b as f64 / n_b.max(1) as f64,
ok_span as f64 / n_span.max(1) as f64,
))
}
pub fn save(varmap: &VarMap, labels_a: &[String], out_dir: &Path) -> Result<(), String> {
std::fs::create_dir_all(out_dir).map_err(|e| e.to_string())?;
varmap.save(out_dir.join("tagger.safetensors")).map_err(|e| format!("save weights: {e}"))?;
let meta = serde_json::json!({
"head_a_labels": labels_a,
"head_b_labels": head_b_labels(),
"ignore_index": IGNORE,
});
std::fs::write(out_dir.join("tagger.json"), serde_json::to_vec_pretty(&meta).map_err(|e| e.to_string())?).map_err(|e| e.to_string())?;
Ok(())
}
fn round4(v: f64) -> f64 {
(v * 10000.0).round() / 10000.0
}
#[derive(Debug, Clone, Serialize, PartialEq)]
pub struct PredictedSpan {
pub start: usize,
pub end: usize,
pub facet: String,
pub text: String,
pub negated: bool,
pub hedged: bool,
pub belief: f32,
}
pub struct TunedTagger {
model: MultiHeadTagger,
tok: Tokenizer,
labels_a: Vec<String>,
device: Device,
max_len: usize,
relations: Option<(crate::relation_train::BiaffineHead, VocabularySpace)>,
}
impl TunedTagger {
pub fn load(dir: &Path, base_dir: &Path, tokenizer: &Path, max_len: usize) -> Result<TunedTagger, String> {
let meta: serde_json::Value =
serde_json::from_slice(&std::fs::read(dir.join("tagger.json")).map_err(|e| format!("tagger.json: {e}"))?)
.map_err(|e| format!("parse tagger.json: {e}"))?;
let labels_a: Vec<String> = meta
.get("head_a_labels")
.and_then(|v| v.as_array())
.map(|a| a.iter().filter_map(|s| s.as_str().map(String::from)).collect())
.ok_or("tagger.json missing head_a_labels")?;
let n_b = head_b_labels().len();
let bert_cfg: BertConfig =
serde_json::from_slice(&std::fs::read(base_dir.join("config.json")).map_err(|e| format!("base config: {e}"))?)
.map_err(|e| format!("parse base config: {e}"))?;
let device = Device::Cpu;
let weights = dir.join("tagger.safetensors");
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[weights.clone()], DType::F32, &device)
.map_err(|e| format!("load {}: {e}", weights.display()))?
};
let model = HrmTagger::new(vb, &hrm_config_from(&bert_cfg), labels_a.len(), n_b).map_err(|e| format!("build model: {e}"))?;
let tok = Tokenizer::from_file(tokenizer).map_err(|e| format!("tokenizer: {e}"))?;
Ok(TunedTagger { model, tok, labels_a, device, max_len, relations: None })
}
pub fn labels(&self) -> &[String] {
&self.labels_a
}
pub fn enable_relations(&mut self, dir: &Path, spec: &VocabularySpace) -> Result<(), String> {
let path = dir.join("relations.safetensors");
if !path.exists() {
return Err(format!("{} not found — run step 5 first", path.display()));
}
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(&[path.clone()], DType::F32, &self.device)
.map_err(|e| format!("load {}: {e}", path.display()))?
};
let hidden = self.model.hidden_size();
let head = crate::relation_train::BiaffineHead::new(vb, 2 * hidden, spec.relation_facets.len() + 1)
.map_err(|e| format!("build head C: {e}"))?;
self.relations = Some((head, spec.clone()));
Ok(())
}
pub fn has_relations(&self) -> bool {
self.relations.is_some()
}
pub fn project(&self, sentence: &str) -> Result<crate::projector::Situation, String> {
use crate::dimensions;
let spans = self.tag(sentence)?;
let mut tokens: Vec<String> = Vec::new();
let mut numbers: Vec<(String, f64)> = Vec::new();
let (mut negated, mut hedged) = (false, false);
for sp in spans.iter().filter(|s| s.facet == "state") {
negated |= sp.negated;
hedged |= sp.hedged;
}
if negated || hedged {
tokens.push(dimensions::state_uri(negated, hedged).to_string());
}
for sp in &spans {
match sp.facet.as_str() {
"state" => {}
"qty" => {
if let Some((field, si)) = crate::units::parse_quantity(&sp.text) {
if let Some(u) = dimensions::qty_uri(&field, si) {
tokens.push(u);
}
numbers.push((field, si));
}
}
"time" => tokens.push(dimensions::time_uri(&sp.text).unwrap_or_else(|| format!("time/{}", crate::projector::slug(&sp.text)))),
"geo" => tokens.push(dimensions::geo_uri(&sp.text)),
facet => tokens.push(dimensions::entity_uri(facet, &sp.text)),
}
}
if let Some((head, spec)) = self.relations.as_ref() {
if let Ok(pairs) = self.bind_relations(sentence, &spans, head, spec) {
for (name, h, t) in pairs {
tokens.push(dimensions::rel_uri(&name, dimensions::Role::Actor));
tokens.push(dimensions::rel_uri(&name, dimensions::Role::Target));
tokens.push(format!("rel/{}/+/{}", crate::projector::slug(&name), dimensions::entity_uri(&h.facet, &h.text)));
tokens.push(format!("rel/{}/-/{}", crate::projector::slug(&name), dimensions::entity_uri(&t.facet, &t.text)));
}
}
}
tokens.sort();
tokens.dedup();
let level = dimensions::belief_level(negated, hedged);
let beliefs = if level == 1.0 { Vec::new() } else { tokens.iter().map(|t| (t.clone(), level)).collect() };
Ok(crate::projector::Situation { tokens, display: vec![sentence.to_string()], numbers, beliefs })
}
pub fn tag(&self, text: &str) -> Result<Vec<PredictedSpan>, String> {
let enc = self.tok.encode(text, true).map_err(|e| format!("tokenize: {e}"))?;
let n = enc.get_ids().len().min(self.max_len);
let ids: Vec<u32> = enc.get_ids()[..n].to_vec();
let offsets: Vec<(usize, usize)> = enc.get_offsets()[..n].to_vec();
let attn = vec![1u32; n];
let ids_t = Tensor::from_vec(ids, (1, n), &self.device).map_err(|e| e.to_string())?;
let attn_t = Tensor::from_vec(attn, (1, n), &self.device).map_err(|e| e.to_string())?;
let (la, lb) = self.model.forward(&ids_t, &attn_t, false).map_err(|e| format!("forward: {e}"))?;
let pa = softmax(&la.i(0).map_err(|e| e.to_string())?, D::Minus1)
.and_then(|t| t.argmax(D::Minus1))
.and_then(|t| t.to_vec1::<u32>())
.map_err(|e| e.to_string())?;
let pb_probs = softmax(&lb.i(0).map_err(|e| e.to_string())?, D::Minus1).map_err(|e| e.to_string())?;
let pb = pb_probs.argmax(D::Minus1).and_then(|t| t.to_vec1::<u32>()).map_err(|e| e.to_string())?;
let pb_dist: Vec<Vec<f32>> = pb_probs.to_vec2::<f32>().map_err(|e| e.to_string())?;
let mut out: Vec<PredictedSpan> = Vec::new();
let mut cur: Option<(String, usize, usize, Vec<u32>, Vec<Vec<f32>>)> = None; let flush = |cur: &mut Option<(String, usize, usize, Vec<u32>, Vec<Vec<f32>>)>, out: &mut Vec<PredictedSpan>, text: &str| {
if let Some((kind, s, e, votes, dists)) = cur.take() {
let ep = if kind.eq_ignore_ascii_case("state") && !dists.is_empty() {
let mut best = (1usize, f32::NEG_INFINITY);
for c in 1..4usize {
let score: f32 = dists.iter().map(|d| d.get(c).copied().unwrap_or(0.0)).sum();
if score > best.1 {
best = (c, score);
}
}
best.0 as u32
} else {
majority(&votes)
};
let (negated, hedged) = match ep {
0 => (false, false),
1 => (false, true),
2 => (true, true),
_ => (true, false),
};
out.push(PredictedSpan {
start: s,
end: e,
facet: kind.to_lowercase(),
text: text[s..e].to_string(),
negated,
hedged,
belief: crate::dimensions::belief_level(negated, hedged),
});
}
};
for i in 0..n {
let (ts, te) = offsets[i];
if te <= ts {
continue; }
let label = self.labels_a.get(pa[i] as usize).map(|s| s.as_str()).unwrap_or("O");
if let Some(kind) = label.strip_prefix("B-") {
flush(&mut cur, &mut out, text);
cur = Some((kind.to_string(), ts, te, vec![pb[i]], vec![pb_dist[i].clone()]));
} else if let Some(kind) = label.strip_prefix("I-") {
match cur.as_mut() {
Some((k, _, e, votes, dists)) if k == kind => {
*e = te;
votes.push(pb[i]);
dists.push(pb_dist[i].clone());
}
_ => flush(&mut cur, &mut out, text), }
} else {
flush(&mut cur, &mut out, text);
}
}
flush(&mut cur, &mut out, text);
Ok(merge_contiguous(snap_to_words(out, text), text))
}
}
fn snap_to_words(spans: Vec<PredictedSpan>, text: &str) -> Vec<PredictedSpan> {
let b = text.as_bytes();
let alnum = |i: usize| -> bool { (b[i] as char).is_alphanumeric() };
spans
.into_iter()
.map(|mut sp| {
while sp.start > 0 && text.is_char_boundary(sp.start - 1) && alnum(sp.start - 1) && alnum(sp.start.min(b.len() - 1)) {
sp.start -= 1;
}
while sp.end < b.len() && text.is_char_boundary(sp.end) && alnum(sp.end) {
sp.end += 1;
}
while sp.end < b.len() && !text.is_char_boundary(sp.end) {
sp.end += 1;
}
sp.text = text[sp.start..sp.end].to_string();
sp
})
.collect()
}
impl TunedTagger {
fn bind_relations(
&self,
sentence: &str,
spans: &[PredictedSpan],
head: &crate::relation_train::BiaffineHead,
spec: &VocabularySpace,
) -> Result<Vec<(String, PredictedSpan, PredictedSpan)>, String> {
use crate::tagger_data::{pair_mask, LabeledSpan};
let ents: Vec<&PredictedSpan> = spans.iter().filter(|s| s.facet != "state").collect();
if ents.len() < 2 || spec.relation_facets.is_empty() {
return Ok(Vec::new());
}
let enc = self.tok.encode(sentence, true).map_err(|e| format!("tokenize: {e}"))?;
let n = enc.get_ids().len().min(self.max_len);
let ids = Tensor::from_vec(enc.get_ids()[..n].to_vec(), (1, n), &self.device).map_err(|e| e.to_string())?;
let attn = Tensor::from_vec(vec![1u32; n], (1, n), &self.device).map_err(|e| e.to_string())?;
let hidden = self.model.hidden(&ids, &attn, false).map_err(|e| format!("encode: {e}"))?;
let offsets: Vec<(usize, usize)> = enc.get_offsets()[..n].to_vec();
let reps: Vec<Tensor> = ents
.iter()
.map(|sp| {
let ls = LabeledSpan { start: sp.start, end: sp.end, facet: sp.facet.clone(), surface: sp.text.clone(), negated: sp.negated, hedged: sp.hedged };
crate::relation_train::pool_span(&hidden, &offsets, &ls)
})
.collect::<candle_core::Result<Vec<_>>>()
.map_err(|e| format!("pool: {e}"))?;
let stacked = Tensor::stack(&reps, 0).map_err(|e| e.to_string())?;
let logits = head.forward(&stacked).map_err(|e| format!("head C: {e}"))?;
let mask = pair_mask(spec);
let n_rel = spec.relation_facets.len() + 1;
let mut out = Vec::new();
for i in 0..ents.len() {
for j in 0..ents.len() {
if i == j {
continue;
}
let row = logits.i((i, j)).map_err(|e| e.to_string())?;
let allow: Vec<f32> = (0..n_rel)
.map(|c| {
if crate::relation_train::type_allowed(spec, &mask, &ents[i].facet, &ents[j].facet, c) {
0.0
} else {
f32::NEG_INFINITY
}
})
.collect();
let allow = Tensor::from_vec(allow, n_rel, &self.device).map_err(|e| e.to_string())?;
let cls = softmax(&(row + allow).map_err(|e| e.to_string())?, D::Minus1)
.and_then(|t| t.argmax(D::Minus1))
.and_then(|t| t.to_scalar::<u32>())
.map_err(|e| e.to_string())? as usize;
if cls > 0 {
if let Some(r) = spec.relation_facets.get(cls - 1) {
out.push((r.name.clone(), ents[i].clone(), ents[j].clone()));
}
}
}
}
Ok(out)
}
}
fn merge_contiguous(spans: Vec<PredictedSpan>, text: &str) -> Vec<PredictedSpan> {
let mut out: Vec<PredictedSpan> = Vec::with_capacity(spans.len());
for sp in spans {
match out.last_mut() {
Some(prev) if prev.facet == sp.facet && prev.end == sp.start => {
prev.end = sp.end;
prev.text = text[prev.start..prev.end].to_string();
prev.negated |= sp.negated;
prev.hedged |= sp.hedged;
prev.belief = crate::dimensions::belief_level(prev.negated, prev.hedged);
}
_ => out.push(sp),
}
}
out
}
fn majority(v: &[u32]) -> u32 {
let mut counts = [0usize; 8];
for &x in v {
if (x as usize) < counts.len() {
counts[x as usize] += 1;
}
}
counts.iter().enumerate().max_by_key(|(_, c)| **c).map(|(i, _)| i as u32).unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tagger_data::{Case, LabeledSpan, RelationLabel};
use crate::vocabulary::{EntityFacet, RelationFacet};
fn spec() -> VocabularySpace {
VocabularySpace {
version: 1,
corpus: "t".into(),
entity_facets: vec![
EntityFacet { name: "org".into(), parent: None, description: "companies".into(), examples: vec![], structural: false },
EntityFacet { name: "system".into(), parent: None, description: "platforms".into(), examples: vec![], structural: false },
],
relation_facets: vec![RelationFacet { name: "develops".into(), head: "org".into(), tail: "system".into() }],
gazetteer: vec![],
metrics: None,
}
}
fn ex(text: &str, a: (usize, usize), b: (usize, usize), negated: bool) -> TaggerExample {
TaggerExample {
text: text.into(),
spans: vec![
LabeledSpan { start: a.0, end: a.1, facet: "org".into(), surface: text[a.0..a.1].into(), negated, hedged: false },
LabeledSpan { start: b.0, end: b.1, facet: "system".into(), surface: text[b.0..b.1].into(), negated, hedged: false },
],
relations: vec![RelationLabel { head: 0, tail: 1, name: "develops".into() }],
case: if negated { Case::Negated } else { Case::Normal },
}
}
fn base_dirs() -> Option<(PathBuf, PathBuf)> {
let home = std::env::var("HOME").ok()?;
let g = |p: &str| glob_first(&format!("{home}/{p}"));
let base = g(".cache/huggingface/hub/models--google--bert_uncased_L-2_H-128_A-2/snapshots/*")?;
let tokdir = g(".cache/huggingface/hub/models--bert-base-uncased/snapshots/*")?;
let tok = tokdir.join("tokenizer.json");
if base.join("model.safetensors").exists() && tok.exists() {
Some((base, tok))
} else {
None
}
}
fn glob_first(pat: &str) -> Option<PathBuf> {
let (dir, _) = pat.rsplit_once('/')?;
std::fs::read_dir(dir).ok()?.filter_map(|e| e.ok()).map(|e| e.path()).find(|p| p.is_dir())
}
#[test]
fn encoding_projects_labels_onto_tokens() {
let Some((_, tok_path)) = base_dirs() else {
eprintln!("skip: no cached bert");
return;
};
let tok = Tokenizer::from_file(&tok_path).unwrap();
let s = spec();
let e = ex("Boeing develops the MQ-28 aircraft.", (0, 6), (20, 25), false);
let enc = encode(&s, &tok, &e, 32).unwrap();
assert_eq!(enc.ids.len(), 32);
assert_eq!(enc.attn.iter().filter(|&&a| a == 1).count(), tok.encode(e.text.as_str(), true).unwrap().get_ids().len());
let labels = head_a_labels(&s);
let named: Vec<&str> = enc.labels_a.iter().filter(|&&l| l != IGNORE && l != 0).map(|&l| labels[l as usize].as_str()).collect();
assert!(named.contains(&"B-ORG"), "got {named:?}");
assert!(named.contains(&"B-SYSTEM"), "got {named:?}");
assert_eq!(enc.labels_a[31], IGNORE);
}
#[test]
fn training_overfits_a_tiny_set() {
let Some((base, tok)) = base_dirs() else {
eprintln!("skip: no cached bert-tiny");
return;
};
let s = spec();
let data = vec![
ex("Boeing develops the MQ-28 aircraft.", (0, 6), (20, 25), false),
ex("Airbus develops the A400M transport.", (0, 6), (20, 25), false),
ex("Saab develops the Gripen fighter jet.", (0, 4), (18, 24), false),
ex("Thales develops the Sonar array system.", (0, 6), (20, 25), false),
ex("Embraer develops the KC-390 airlifter.", (0, 7), (21, 27), false),
ex("Lockheed does not develop the F-35 jet.", (0, 8), (30, 34), true),
];
let cfg = TrainConfig { base_dir: base, tokenizer: tok, epochs: 60, lr: 3e-3, batch: 6, max_len: 32, ..Default::default() };
let (varmap, rep, labels) = train(&s, &data, &cfg).expect("train");
eprintln!("{}", serde_json::to_string_pretty(&rep).unwrap());
let expect = 1 + 2 * (s.taggable_facets().len() + crate::tagger_data::STRUCTURAL_KINDS.len());
assert_eq!(labels.len(), expect);
let first = rep.epochs.first().unwrap().total;
let last = rep.epochs.last().unwrap().total;
assert!(last < first, "loss must decrease: {first} → {last}");
assert!(rep.train_acc_a > 0.8, "Head A should overfit 3 examples (got {})", rep.train_acc_a);
assert!(rep.train_acc_b > 0.8, "Head B should overfit 3 examples (got {})", rep.train_acc_b);
let dir = std::env::temp_dir().join(format!("steeldb-tagger-{}", std::process::id()));
save(&varmap, &labels, &dir).unwrap();
assert!(dir.join("tagger.safetensors").exists() && dir.join("tagger.json").exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn tuned_tagger_round_trips_and_tags() {
let Some((base, tok)) = base_dirs() else {
eprintln!("skip: no cached bert-tiny");
return;
};
let s = spec();
let data = vec![
ex("Boeing develops the MQ-28 aircraft.", (0, 6), (20, 25), false),
ex("Airbus develops the A400M transport.", (0, 6), (20, 25), false),
ex("Saab develops the Gripen fighter jet.", (0, 4), (18, 24), false),
ex("Thales develops the Sonar array system.", (0, 6), (20, 25), false),
ex("Embraer develops the KC-390 airlifter.", (0, 7), (21, 27), false),
ex("Lockheed does not develop the F-35 jet.", (0, 8), (30, 34), true),
];
let cfg = TrainConfig { base_dir: base.clone(), tokenizer: tok.clone(), epochs: 60, lr: 3e-3, batch: 6, max_len: 32, ..Default::default() };
let (varmap, _rep, labels) = train(&s, &data, &cfg).expect("train");
let dir = std::env::temp_dir().join(format!("steeldb-tagger-rt-{}", std::process::id()));
save(&varmap, &labels, &dir).unwrap();
let tt = TunedTagger::load(&dir, &base, &tok, 32).expect("load tuned");
assert_eq!(tt.labels().len(), labels.len());
let spans = tt.tag("Boeing develops the MQ-28 aircraft.").expect("tag");
eprintln!("TAGGED: {spans:?}");
assert!(spans.iter().any(|p| p.facet == "org" && p.text.contains("Boeing")), "got {spans:?}");
assert!(spans.iter().any(|p| p.facet == "system"), "got {spans:?}");
let neg = tt.tag("Lockheed does not develop the F-35 jet.").expect("tag");
assert!(neg.iter().any(|p| p.negated && p.belief < 0.0), "negation must set belief<0: {neg:?}");
for w in spans.windows(2) {
assert!(w[0].end <= w[1].start, "spans must not overlap: {spans:?}");
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn projection_reaches_the_bitmap_with_polarity() {
let Some((base, tok)) = base_dirs() else {
eprintln!("skip: no cached bert-tiny");
return;
};
let s = spec();
let data = vec![
ex("Boeing develops the MQ-28 aircraft.", (0, 6), (20, 25), false),
ex("Airbus develops the A400M transport.", (0, 6), (20, 25), false),
ex("Saab develops the Gripen fighter jet.", (0, 4), (18, 24), false),
ex("Thales develops the Sonar array system.", (0, 6), (20, 25), false),
ex("Embraer develops the KC-390 airlifter.", (0, 7), (21, 27), false),
ex("Lockheed does not develop the F-35 jet.", (0, 8), (30, 34), true),
];
let cfg = TrainConfig { base_dir: base.clone(), tokenizer: tok.clone(), epochs: 60, lr: 3e-3, batch: 6, max_len: 32, ..Default::default() };
let (varmap, _r, labels) = train(&s, &data, &cfg).expect("train");
let dir = std::env::temp_dir().join(format!("steeldb-proj-{}", std::process::id()));
save(&varmap, &labels, &dir).unwrap();
let tt = TunedTagger::load(&dir, &base, &tok, 32).expect("load");
let sit = tt.project("Boeing develops the MQ-28 aircraft.").expect("project");
eprintln!("PROJECTED: {:?}", sit.tokens);
assert!(sit.tokens.iter().any(|t| t.starts_with("org/")), "typed entity token expected: {:?}", sit.tokens);
assert!(sit.beliefs.is_empty(), "asserted sentence carries default +1 polarity");
let mut corpus = crate::db::Corpus::new_incremental("t", vec!["text".into()], crate::projector::CorpusKind::Csv);
corpus.add_situation_polar(sit.tokens.clone(), sit.display.clone(), sit.numbers.clone(), sit.beliefs.clone());
let hits = corpus.query("org/*", 10);
assert_eq!(hits.count, 1, "facet wildcard must match the typed token");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn state_spans_never_decode_as_asserted() {
let Some((base, tok)) = base_dirs() else {
eprintln!("skip: no cached bert-tiny");
return;
};
let s = spec();
let mut neg = ex("Lockheed does not develop the F-35 jet.", (0, 8), (30, 34), true);
neg.spans.push(LabeledSpan { start: 9, end: 17, facet: "state".into(), surface: "does not".into(), negated: true, hedged: false });
neg.spans.sort_by_key(|x| x.start);
let mut hedge = ex("Boeing may develop the MQ-28 aircraft.", (0, 6), (23, 28), false);
hedge.spans.push(LabeledSpan { start: 7, end: 10, facet: "state".into(), surface: "may".into(), negated: false, hedged: true });
hedge.spans.iter_mut().for_each(|x| { if x.facet != "state" { x.hedged = true; } });
hedge.spans.sort_by_key(|x| x.start);
let data = vec![neg.clone(), hedge.clone(), neg.clone(), hedge.clone(), neg.clone(), hedge];
let cfg = TrainConfig { base_dir: base.clone(), tokenizer: tok.clone(), epochs: 60, lr: 3e-3, batch: 6, max_len: 32, ..Default::default() };
let (varmap, _r, labels) = train(&s, &data, &cfg).expect("train");
let dir = std::env::temp_dir().join(format!("steeldb-state-{}", std::process::id()));
save(&varmap, &labels, &dir).unwrap();
let tt = TunedTagger::load(&dir, &base, &tok, 32).expect("load");
for text in ["Lockheed does not develop the F-35 jet.", "Boeing may develop the MQ-28 aircraft."] {
for sp in tt.tag(text).expect("tag").iter().filter(|s| s.facet == "state") {
assert!(sp.negated || sp.hedged, "a state cue decoded as asserted: {sp:?}");
assert!(sp.belief < 1.0, "asserted belief on a cue span: {sp:?}");
}
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn snapping_completes_clipped_words() {
let text = "Thales supplies the Aegis system.";
let p = |s: usize, e: usize, f: &str| PredictedSpan {
start: s, end: e, facet: f.into(), text: text[s..e].into(), negated: false, hedged: false, belief: 1.0,
};
let out = super::snap_to_words(vec![p(0, 3, "org")], text);
assert_eq!(out[0].text, "Thales");
let out2 = super::snap_to_words(vec![p(20, 22, "system")], text);
assert_eq!(out2[0].text, "Aegis");
let out3 = super::snap_to_words(vec![p(0, 6, "org")], text);
assert_eq!(out3[0].text, "Thales");
assert_eq!(out3[0].end, 6);
}
#[test]
fn contiguous_subword_spans_merge_but_distinct_entities_do_not() {
let text = "Aegis and Boeing Airbus";
let p = |s: usize, e: usize, f: &str| PredictedSpan {
start: s, end: e, facet: f.into(), text: text[s..e].into(), negated: false, hedged: false, belief: 1.0,
};
let merged = super::merge_contiguous(vec![p(0, 2, "system"), p(2, 5, "system")], text);
assert_eq!(merged.len(), 1);
assert_eq!(merged[0].text, "Aegis");
let kept = super::merge_contiguous(vec![p(10, 16, "org"), p(17, 23, "org")], text);
assert_eq!(kept.len(), 2);
let mut a = p(0, 2, "system");
let mut b = p(2, 5, "system");
b.negated = true;
a.belief = 1.0;
let m2 = super::merge_contiguous(vec![a, b], text);
assert_eq!(m2.len(), 1);
assert!(m2[0].negated && m2[0].belief < 0.0);
}
}