use burn::module::{Module, Param, ParamId};
use burn::optim::{AdamWConfig, GradientsParams, Optimizer};
use burn::prelude::*;
use burn::tensor::activation;
use burn::tensor::backend::AutodiffBackend;
#[derive(Module, Debug)]
pub struct BurnComplEx<B: Backend> {
entity_re: Param<Tensor<B, 2>>,
entity_im: Param<Tensor<B, 2>>,
relation_re: Param<Tensor<B, 2>>,
relation_im: Param<Tensor<B, 2>>,
}
#[derive(Debug, Clone)]
pub struct BurnTrainConfig {
pub dim: usize,
pub init_scale: f64,
pub lr: f64,
pub label_smoothing: f64,
pub n3_reg: f64,
pub batch_size: usize,
pub epochs: usize,
pub log_interval: usize,
}
impl Default for BurnTrainConfig {
fn default() -> Self {
Self {
dim: 200,
init_scale: 1e-3,
lr: 0.001,
label_smoothing: 0.1,
n3_reg: 0.0,
batch_size: 512,
epochs: 100,
log_interval: 10,
}
}
}
pub struct BurnTrainResult {
pub entity_vecs: Vec<Vec<f32>>,
pub relation_vecs: Vec<Vec<f32>>,
pub dim: usize,
pub losses: Vec<f32>,
}
impl BurnTrainResult {
pub fn to_complex(&self) -> crate::ComplEx {
crate::ComplEx::from_vecs(
self.entity_vecs.clone(),
self.relation_vecs.clone(),
self.dim,
)
}
}
fn init_model<B: AutodiffBackend>(
num_entities: usize,
num_relations: usize,
dim: usize,
init_scale: f64,
device: &B::Device,
) -> BurnComplEx<B> {
let mk = |rows, cols| {
Param::initialized(
ParamId::new(),
Tensor::<B, 2>::random(
[rows, cols],
burn::tensor::Distribution::Normal(0.0, init_scale),
device,
)
.require_grad(),
)
};
BurnComplEx {
entity_re: mk(num_entities, dim),
entity_im: mk(num_entities, dim),
relation_re: mk(num_relations, dim),
relation_im: mk(num_relations, dim),
}
}
fn score_1n<B: Backend>(
model: &BurnComplEx<B>,
heads: &Tensor<B, 1, Int>,
rels: &Tensor<B, 1, Int>,
) -> Tensor<B, 2> {
let h_re = model.entity_re.val().select(0, heads.clone());
let h_im = model.entity_im.val().select(0, heads.clone());
let r_re = model.relation_re.val().select(0, rels.clone());
let r_im = model.relation_im.val().select(0, rels.clone());
let hr_re = h_re.clone() * r_re.clone() - h_im.clone() * r_im.clone();
let hr_im = h_re * r_im + h_im * r_re;
let e_re = model.entity_re.val();
let e_im = model.entity_im.val();
hr_re.matmul(e_re.transpose()) + hr_im.matmul(e_im.transpose())
}
fn score_1n_heads<B: Backend>(
model: &BurnComplEx<B>,
rels: &Tensor<B, 1, Int>,
tails: &Tensor<B, 1, Int>,
) -> Tensor<B, 2> {
let r_re = model.relation_re.val().select(0, rels.clone());
let r_im = model.relation_im.val().select(0, rels.clone());
let t_re = model.entity_re.val().select(0, tails.clone());
let t_im = model.entity_im.val().select(0, tails.clone());
let rc_re = r_re.clone() * t_re.clone() + r_im.clone() * t_im.clone();
let rc_im = r_im * t_re - r_re * t_im;
let e_re = model.entity_re.val();
let e_im = model.entity_im.val();
rc_re.matmul(e_re.transpose()) - rc_im.matmul(e_im.transpose())
}
pub fn train_complex<B: AutodiffBackend>(
train_triples: &[crate::dataset::TripleIds],
num_entities: usize,
num_relations: usize,
config: &BurnTrainConfig,
device: &B::Device,
) -> BurnTrainResult {
let mut model = init_model::<B>(
num_entities,
num_relations,
config.dim,
config.init_scale,
device,
);
let mut optim = AdamWConfig::new()
.with_epsilon(1e-8)
.with_weight_decay(0.0)
.init::<B, BurnComplEx<B>>();
let n_triples = train_triples.len();
let batch_size = config.batch_size.min(n_triples);
let eps = config.label_smoothing;
let mut losses = Vec::with_capacity(config.epochs);
let mut indices: Vec<usize> = (0..n_triples).collect();
for epoch in 0..config.epochs {
let epoch_start = std::time::Instant::now();
{
use rand::seq::SliceRandom;
indices.shuffle(&mut rand::rng());
}
let mut epoch_loss = 0.0_f64;
let mut n_batches = 0u32;
let mut offset = 0;
while offset < n_triples {
let end = (offset + batch_size).min(n_triples);
let batch_idx = &indices[offset..end];
let actual_bs = batch_idx.len();
offset = end;
let heads_data: Vec<i64> = batch_idx
.iter()
.map(|&i| train_triples[i].head as i64)
.collect();
let rels_data: Vec<i64> = batch_idx
.iter()
.map(|&i| train_triples[i].relation as i64)
.collect();
let tails_data: Vec<i64> = batch_idx
.iter()
.map(|&i| train_triples[i].tail as i64)
.collect();
let heads = Tensor::<B, 1, Int>::from_data(
burn::tensor::TensorData::new(heads_data, [actual_bs]),
device,
);
let rels = Tensor::<B, 1, Int>::from_data(
burn::tensor::TensorData::new(rels_data, [actual_bs]),
device,
);
let tails = Tensor::<B, 1, Int>::from_data(
burn::tensor::TensorData::new(tails_data.clone(), [actual_bs]),
device,
);
let current = model.clone();
let tail_scores = score_1n(¤t, &heads, &rels);
let tail_log_probs = activation::log_softmax(tail_scores, 1);
let head_scores = score_1n_heads(¤t, &rels, &tails);
let head_log_probs = activation::log_softmax(head_scores, 1);
let tail_ids = tails.clone().unsqueeze_dim(1); let t_nll = tail_log_probs
.clone()
.gather(1, tail_ids)
.squeeze::<1>()
.neg()
.mean();
let head_ids = heads.clone().unsqueeze_dim(1); let h_nll = head_log_probs
.clone()
.gather(1, head_ids)
.squeeze::<1>()
.neg()
.mean();
let nll = (t_nll + h_nll) / 2.0;
let loss = if eps > 0.0 {
let tail_uniform = tail_log_probs.mean().neg();
let head_uniform = head_log_probs.mean().neg();
let uniform = (tail_uniform + head_uniform) / 2.0;
nll * (1.0 - eps) + uniform * eps
} else {
nll
};
let loss_val: f32 = loss.clone().inner().into_scalar().to_f32();
let grads = GradientsParams::from_grads(loss.backward(), ¤t);
if loss_val.is_finite() {
model = optim.step(config.lr, current, grads);
}
epoch_loss += loss_val as f64;
n_batches += 1;
}
let avg_loss = (epoch_loss / n_batches as f64) as f32;
losses.push(avg_loss);
if config.log_interval > 0 && (epoch + 1) % config.log_interval == 0 {
eprintln!(
"epoch {:>4} | loss {:.4} | {:.1}s",
epoch + 1,
avg_loss,
epoch_start.elapsed().as_secs_f32(),
);
}
}
let dim = config.dim;
let extract = |re: &Param<Tensor<B, 2>>, im: &Param<Tensor<B, 2>>| -> Vec<Vec<f32>> {
let re_data: Vec<f32> = re.val().into_data().to_vec().unwrap();
let im_data: Vec<f32> = im.val().into_data().to_vec().unwrap();
let n = re_data.len() / dim;
(0..n)
.map(|i| {
let mut v = Vec::with_capacity(dim * 2);
v.extend_from_slice(&re_data[i * dim..(i + 1) * dim]);
v.extend_from_slice(&im_data[i * dim..(i + 1) * dim]);
v
})
.collect()
};
BurnTrainResult {
entity_vecs: extract(&model.entity_re, &model.entity_im),
relation_vecs: extract(&model.relation_re, &model.relation_im),
dim,
losses,
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BurnModelType {
TransE,
RotatE,
ComplEx,
DistMult,
}
impl BurnModelType {
fn entity_dim(self, dim: usize) -> usize {
match self {
BurnModelType::ComplEx | BurnModelType::RotatE => 2 * dim,
BurnModelType::TransE | BurnModelType::DistMult => dim,
}
}
fn relation_dim(self, dim: usize) -> usize {
match self {
BurnModelType::ComplEx => 2 * dim,
_ => dim,
}
}
}
#[derive(Module, Debug)]
pub struct BurnKge<B: Backend> {
entity: Param<Tensor<B, 2>>,
relation: Param<Tensor<B, 2>>,
}
pub struct BurnKgeResult {
pub entity_vecs: Vec<Vec<f32>>,
pub relation_vecs: Vec<Vec<f32>>,
pub dim: usize,
pub model_type: BurnModelType,
pub losses: Vec<f32>,
}
impl BurnKgeResult {
pub fn to_scorer(&self) -> Box<dyn crate::Scorer + Sync> {
let e = self.entity_vecs.clone();
let r = self.relation_vecs.clone();
match self.model_type {
BurnModelType::TransE => Box::new(crate::TransE::from_vecs(e, r, self.dim)),
BurnModelType::DistMult => Box::new(crate::DistMult::from_vecs(e, r, self.dim)),
BurnModelType::ComplEx => Box::new(crate::ComplEx::from_vecs(e, r, self.dim)),
BurnModelType::RotatE => Box::new(crate::RotatE::from_vecs(e, r, self.dim, 12.0)),
}
}
}
fn re_im<B: Backend>(t: Tensor<B, 2>, dim: usize) -> (Tensor<B, 2>, Tensor<B, 2>) {
let rows = t.dims()[0];
(
t.clone().slice([0..rows, 0..dim]),
t.slice([0..rows, dim..2 * dim]),
)
}
fn neg_sq_dist<B: Backend>(hr: Tensor<B, 2>, ent: Tensor<B, 2>) -> Tensor<B, 2> {
let hr_sq = hr.clone().powf_scalar(2.0).sum_dim(1); let ent_sq = ent.clone().powf_scalar(2.0).sum_dim(1).transpose(); let cross = hr.matmul(ent.transpose()); (hr_sq + ent_sq - cross.mul_scalar(2.0)).neg()
}
fn score_1n_kge<B: Backend>(
model: &BurnKge<B>,
mt: BurnModelType,
dim: usize,
heads: &Tensor<B, 1, Int>,
rels: &Tensor<B, 1, Int>,
) -> Tensor<B, 2> {
let h = model.entity.val().select(0, heads.clone());
let r = model.relation.val().select(0, rels.clone());
let ent = model.entity.val();
match mt {
BurnModelType::TransE => neg_sq_dist(h + r, ent),
BurnModelType::DistMult => (h * r).matmul(ent.transpose()),
BurnModelType::ComplEx => {
let (h_re, h_im) = re_im(h, dim);
let (r_re, r_im) = re_im(r, dim);
let hr_re = h_re.clone() * r_re.clone() - h_im.clone() * r_im.clone();
let hr_im = h_re * r_im + h_im * r_re;
let (e_re, e_im) = re_im(ent, dim);
hr_re.matmul(e_re.transpose()) + hr_im.matmul(e_im.transpose())
}
BurnModelType::RotatE => {
let (h_re, h_im) = re_im(h, dim);
let cos = r.clone().cos();
let sin = r.sin();
let hr_re = h_re.clone() * cos.clone() - h_im.clone() * sin.clone();
let hr_im = h_re * sin + h_im * cos;
neg_sq_dist(Tensor::cat(vec![hr_re, hr_im], 1), ent)
}
}
}
fn score_1n_heads_kge<B: Backend>(
model: &BurnKge<B>,
mt: BurnModelType,
dim: usize,
rels: &Tensor<B, 1, Int>,
tails: &Tensor<B, 1, Int>,
) -> Tensor<B, 2> {
let r = model.relation.val().select(0, rels.clone());
let t = model.entity.val().select(0, tails.clone());
let ent = model.entity.val();
match mt {
BurnModelType::TransE => neg_sq_dist(t - r, ent),
BurnModelType::DistMult => (r * t).matmul(ent.transpose()),
BurnModelType::ComplEx => {
let (r_re, r_im) = re_im(r, dim);
let (t_re, t_im) = re_im(t, dim);
let rc_re = r_re.clone() * t_re.clone() + r_im.clone() * t_im.clone();
let rc_im = r_im * t_re - r_re * t_im;
let (e_re, e_im) = re_im(ent, dim);
rc_re.matmul(e_re.transpose()) - rc_im.matmul(e_im.transpose())
}
BurnModelType::RotatE => {
let (t_re, t_im) = re_im(t, dim);
let cos = r.clone().cos();
let sin = r.sin();
let tr_re = t_re.clone() * cos.clone() + t_im.clone() * sin.clone();
let tr_im = t_im * cos - t_re * sin;
neg_sq_dist(Tensor::cat(vec![tr_re, tr_im], 1), ent)
}
}
}
pub fn train_kge<B: AutodiffBackend>(
train_triples: &[crate::dataset::TripleIds],
num_entities: usize,
num_relations: usize,
model_type: BurnModelType,
config: &BurnTrainConfig,
device: &B::Device,
) -> BurnKgeResult {
let ent_dim = model_type.entity_dim(config.dim);
let rel_dim = model_type.relation_dim(config.dim);
let mk = |rows, cols| {
Param::initialized(
ParamId::new(),
Tensor::<B, 2>::random(
[rows, cols],
burn::tensor::Distribution::Normal(0.0, config.init_scale),
device,
)
.require_grad(),
)
};
let mut model = BurnKge {
entity: mk(num_entities, ent_dim),
relation: mk(num_relations, rel_dim),
};
let mut optim = AdamWConfig::new()
.with_epsilon(1e-8)
.with_weight_decay(0.0)
.init::<B, BurnKge<B>>();
let n_triples = train_triples.len();
let batch_size = config.batch_size.min(n_triples).max(1);
let eps = config.label_smoothing;
let mut losses = Vec::with_capacity(config.epochs);
let mut indices: Vec<usize> = (0..n_triples).collect();
for _epoch in 0..config.epochs {
{
use rand::seq::SliceRandom;
indices.shuffle(&mut rand::rng());
}
let mut epoch_loss_acc: Option<Tensor<<B as AutodiffBackend>::InnerBackend, 1>> = None;
let mut n_batches = 0u32;
let mut offset = 0;
while offset < n_triples {
let end = (offset + batch_size).min(n_triples);
let batch_idx = &indices[offset..end];
let bs = batch_idx.len();
offset = end;
let hd: Vec<i64> = batch_idx
.iter()
.map(|&i| train_triples[i].head as i64)
.collect();
let rd: Vec<i64> = batch_idx
.iter()
.map(|&i| train_triples[i].relation as i64)
.collect();
let tdv: Vec<i64> = batch_idx
.iter()
.map(|&i| train_triples[i].tail as i64)
.collect();
let heads =
Tensor::<B, 1, Int>::from_data(burn::tensor::TensorData::new(hd, [bs]), device);
let rels =
Tensor::<B, 1, Int>::from_data(burn::tensor::TensorData::new(rd, [bs]), device);
let tails =
Tensor::<B, 1, Int>::from_data(burn::tensor::TensorData::new(tdv, [bs]), device);
let current = model.clone();
let tail_scores = score_1n_kge(¤t, model_type, config.dim, &heads, &rels);
let tail_lp = activation::log_softmax(tail_scores, 1);
let head_scores = score_1n_heads_kge(¤t, model_type, config.dim, &rels, &tails);
let head_lp = activation::log_softmax(head_scores, 1);
let t_nll = tail_lp
.clone()
.gather(1, tails.clone().unsqueeze_dim(1))
.squeeze::<1>()
.neg()
.mean();
let h_nll = head_lp
.clone()
.gather(1, heads.clone().unsqueeze_dim(1))
.squeeze::<1>()
.neg()
.mean();
let nll = (t_nll + h_nll) / 2.0;
let loss = if eps > 0.0 {
let uniform = (tail_lp.mean().neg() + head_lp.mean().neg()) / 2.0;
nll * (1.0 - eps) + uniform * eps
} else {
nll
};
let loss_inner = loss.clone().inner();
let grads = GradientsParams::from_grads(loss.backward(), ¤t);
model = optim.step(config.lr, current, grads);
epoch_loss_acc = Some(match epoch_loss_acc {
Some(acc) => acc + loss_inner,
None => loss_inner,
});
n_batches += 1;
}
let avg = match epoch_loss_acc {
Some(acc) => acc.into_scalar().to_f32() / n_batches.max(1) as f32,
None => 0.0,
};
losses.push(avg);
}
let extract = |p: &Param<Tensor<B, 2>>, cols: usize| -> Vec<Vec<f32>> {
let data: Vec<f32> = p.val().into_data().to_vec().unwrap();
data.chunks(cols).map(<[f32]>::to_vec).collect()
};
BurnKgeResult {
entity_vecs: extract(&model.entity, ent_dim),
relation_vecs: extract(&model.relation, rel_dim),
dim: config.dim,
model_type,
losses,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dataset::TripleIds;
use crate::Scorer;
fn tid(h: usize, r: usize, t: usize) -> TripleIds {
TripleIds::new(h, r, t)
}
#[cfg(feature = "burn-ndarray")]
type TestBackend = burn::backend::Autodiff<burn_ndarray::NdArray>;
#[cfg(feature = "burn-ndarray")]
fn test_device() -> <TestBackend as Backend>::Device {
burn_ndarray::NdArrayDevice::Cpu
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_complex_smoke() {
let triples = vec![tid(0, 0, 1), tid(1, 0, 2), tid(2, 1, 0), tid(0, 1, 2)];
let config = BurnTrainConfig {
dim: 8,
epochs: 10,
batch_size: 4,
..BurnTrainConfig::default()
};
let result = train_complex::<TestBackend>(&triples, 3, 2, &config, &test_device());
assert_eq!(result.losses.len(), 10);
assert!(result.losses.iter().all(|l| l.is_finite()));
let model = result.to_complex();
assert_eq!(model.num_entities(), 3);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_complex_loss_decreases() {
let triples: Vec<_> = (0..20).map(|i| tid(i % 5, i % 2, (i + 1) % 5)).collect();
let config = BurnTrainConfig {
dim: 16,
epochs: 30,
batch_size: 10,
lr: 0.001,
..BurnTrainConfig::default()
};
let result = train_complex::<TestBackend>(&triples, 5, 2, &config, &test_device());
let first = result.losses[0];
let last = *result.losses.last().unwrap();
assert!(
last < first,
"Burn ComplEx loss should decrease: {first} -> {last}"
);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_complex_scores_match_cpu_reference() {
let (ne, nr, dim) = (5usize, 2usize, 4usize);
let val = |i: usize| ((i * 37 + 11) % 19) as f32 / 19.0 - 0.5;
let ent_rows: Vec<Vec<f32>> = (0..ne)
.map(|e| (0..dim * 2).map(|j| val(e * 31 + j)).collect())
.collect();
let rel_rows: Vec<Vec<f32>> = (0..nr)
.map(|r| (0..dim * 2).map(|j| val(r * 53 + j + 7)).collect())
.collect();
let cpu = crate::ComplEx::from_vecs(ent_rows.clone(), rel_rows.clone(), dim);
let device = test_device();
let param = |rows: &[Vec<f32>], lo: usize, hi: usize| {
let flat: Vec<f32> = rows
.iter()
.flat_map(|r| r[lo..hi].iter().copied())
.collect();
Param::initialized(
ParamId::new(),
Tensor::<TestBackend, 2>::from_data(
burn::tensor::TensorData::new(flat, [rows.len(), hi - lo]),
&device,
),
)
};
let model = BurnComplEx::<TestBackend> {
entity_re: param(&ent_rows, 0, dim),
entity_im: param(&ent_rows, dim, dim * 2),
relation_re: param(&rel_rows, 0, dim),
relation_im: param(&rel_rows, dim, dim * 2),
};
let kge = BurnKge::<TestBackend> {
entity: param(&ent_rows, 0, dim * 2),
relation: param(&rel_rows, 0, dim * 2),
};
let scores =
|t: Tensor<TestBackend, 2>| -> Vec<f32> { t.into_data().to_vec::<f32>().unwrap() };
for r in 0..nr {
for x in 0..ne {
let ids = |i: usize| {
Tensor::<TestBackend, 1, Int>::from_data(
burn::tensor::TensorData::new(vec![i as i64], [1]),
&device,
)
};
let (rels, xs) = (ids(r), ids(x));
let mt = BurnModelType::ComplEx;
let cases = [
(
"heads",
scores(score_1n_heads(&model, &rels, &xs)),
cpu.score_all_heads(r, x),
),
(
"tails",
scores(score_1n(&model, &xs, &rels)),
cpu.score_all_tails(x, r),
),
(
"kge heads",
scores(score_1n_heads_kge(&kge, mt, dim, &rels, &xs)),
cpu.score_all_heads(r, x),
),
(
"kge tails",
scores(score_1n_kge(&kge, mt, dim, &xs, &rels)),
cpu.score_all_tails(x, r),
),
];
for (name, burn_scores, cpu_energies) in cases {
assert_eq!(burn_scores.len(), ne);
for (e, (b, c)) in burn_scores.iter().zip(cpu_energies.iter()).enumerate() {
assert!(
(b + c).abs() < 1e-4,
"{name} mismatch at rel={r} x={x} entity={e}: burn={b} cpu={}",
-c
);
}
}
}
}
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_distmult_scores_match_cpu_reference() {
let (ne, nr, dim) = (5usize, 2usize, 4usize);
let val = |i: usize| ((i * 37 + 11) % 19) as f32 / 19.0 - 0.5;
let ent_rows: Vec<Vec<f32>> = (0..ne)
.map(|e| (0..dim).map(|j| val(e * 31 + j)).collect())
.collect();
let rel_rows: Vec<Vec<f32>> = (0..nr)
.map(|r| (0..dim).map(|j| val(r * 53 + j + 7)).collect())
.collect();
let cpu = crate::DistMult::from_vecs(ent_rows.clone(), rel_rows.clone(), dim);
let device = test_device();
let param = |rows: &[Vec<f32>]| {
let cols = rows[0].len();
let flat: Vec<f32> = rows.iter().flat_map(|r| r.iter().copied()).collect();
Param::initialized(
ParamId::new(),
Tensor::<TestBackend, 2>::from_data(
burn::tensor::TensorData::new(flat, [rows.len(), cols]),
&device,
),
)
};
let kge = BurnKge::<TestBackend> {
entity: param(&ent_rows),
relation: param(&rel_rows),
};
let scores =
|t: Tensor<TestBackend, 2>| -> Vec<f32> { t.into_data().to_vec::<f32>().unwrap() };
let mt = BurnModelType::DistMult;
for r in 0..nr {
for x in 0..ne {
let ids = |i: usize| {
Tensor::<TestBackend, 1, Int>::from_data(
burn::tensor::TensorData::new(vec![i as i64], [1]),
&device,
)
};
let (rels, xs) = (ids(r), ids(x));
let cases = [
(
"kge heads",
scores(score_1n_heads_kge(&kge, mt, dim, &rels, &xs)),
cpu.score_all_heads(r, x),
),
(
"kge tails",
scores(score_1n_kge(&kge, mt, dim, &xs, &rels)),
cpu.score_all_tails(x, r),
),
];
for (name, burn_scores, cpu_energies) in cases {
assert_eq!(burn_scores.len(), ne);
for (e, (b, c)) in burn_scores.iter().zip(cpu_energies.iter()).enumerate() {
assert!(
(b + c).abs() < 1e-4,
"{name} mismatch at rel={r} x={x} entity={e}: burn={b} cpu={}",
-c
);
}
}
}
}
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_complex_achieves_nonzero_mrr() {
let triples = vec![tid(0, 0, 1), tid(1, 0, 2), tid(2, 0, 3), tid(3, 0, 4)];
let config = BurnTrainConfig {
dim: 32,
epochs: 200,
batch_size: 4,
lr: 0.001,
..BurnTrainConfig::default()
};
let result = train_complex::<TestBackend>(&triples, 5, 1, &config, &test_device());
let model = result.to_complex();
let ds = crate::dataset::Dataset::new(
triples
.iter()
.map(|t| {
crate::dataset::Triple::new(
t.head.to_string(),
t.relation.to_string(),
t.tail.to_string(),
)
})
.collect(),
Vec::new(),
Vec::new(),
)
.into_interned();
let filter = crate::dataset::FilterIndex::from_dataset(&ds);
let metrics = crate::eval::evaluate_link_prediction(&model, &triples, &filter, 5);
assert!(
metrics.mrr > 0.3,
"Burn ComplEx should achieve MRR > 0.3, got {:.4}",
metrics.mrr
);
}
#[cfg(feature = "burn-ndarray")]
fn kge_loss_decreases(mt: BurnModelType) {
let triples: Vec<_> = (0..20).map(|i| tid(i % 5, i % 2, (i + 1) % 5)).collect();
let config = BurnTrainConfig {
dim: 16,
epochs: 40,
batch_size: 10,
lr: 0.005,
..BurnTrainConfig::default()
};
let result = train_kge::<TestBackend>(&triples, 5, 2, mt, &config, &test_device());
assert_eq!(result.losses.len(), 40);
assert!(
result.losses.iter().all(|l| l.is_finite()),
"{mt:?} produced a non-finite loss"
);
let (first, last) = (result.losses[0], *result.losses.last().unwrap());
assert!(
last < first,
"{mt:?} loss should decrease: {first} -> {last}"
);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_transe_loss_decreases() {
kge_loss_decreases(BurnModelType::TransE);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_distmult_loss_decreases() {
kge_loss_decreases(BurnModelType::DistMult);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_rotate_loss_decreases() {
kge_loss_decreases(BurnModelType::RotatE);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_distmult_achieves_nonzero_mrr() {
let triples = vec![tid(0, 0, 1), tid(1, 0, 2), tid(2, 0, 3), tid(3, 0, 4)];
let config = BurnTrainConfig {
dim: 32,
epochs: 200,
batch_size: 4,
lr: 0.01,
..BurnTrainConfig::default()
};
let result = train_kge::<TestBackend>(
&triples,
5,
1,
BurnModelType::DistMult,
&config,
&test_device(),
);
let model = crate::DistMult::from_vecs(
result.entity_vecs.clone(),
result.relation_vecs.clone(),
result.dim,
);
let ds = crate::dataset::Dataset::new(
triples
.iter()
.map(|t| {
crate::dataset::Triple::new(
t.head.to_string(),
t.relation.to_string(),
t.tail.to_string(),
)
})
.collect(),
Vec::new(),
Vec::new(),
)
.into_interned();
let filter = crate::dataset::FilterIndex::from_dataset(&ds);
let metrics = crate::eval::evaluate_link_prediction(&model, &triples, &filter, 5);
assert!(
metrics.mrr > 0.3,
"Burn DistMult should achieve MRR > 0.3, got {:.4}",
metrics.mrr
);
}
#[cfg(feature = "burn-ndarray")]
fn kge_ranks_match_cpu_reference(mt: BurnModelType) {
let (ne, nr, dim) = (6usize, 2usize, 4usize);
let val = |i: usize| {
let x = ((i as f64) * 0.618_033_988_749_895).fract();
(2.0 * x - 1.0) as f32
};
let (ent_w, rel_w) = match mt {
BurnModelType::TransE => (dim, dim),
BurnModelType::RotatE => (2 * dim, dim),
other => panic!("rank-parity helper is for distance models, not {other:?}"),
};
let ent_rows: Vec<Vec<f32>> = (0..ne)
.map(|e| (0..ent_w).map(|j| val(e * 31 + j + 1)).collect())
.collect();
let rel_rows: Vec<Vec<f32>> = (0..nr)
.map(|r| (0..rel_w).map(|j| val(r * 53 + j + 7)).collect())
.collect();
let cpu: Box<dyn Scorer + Sync> = match mt {
BurnModelType::TransE => Box::new(crate::TransE::from_vecs(
ent_rows.clone(),
rel_rows.clone(),
dim,
)),
BurnModelType::RotatE => Box::new(crate::RotatE::from_vecs(
ent_rows.clone(),
rel_rows.clone(),
dim,
12.0,
)),
other => panic!("rank-parity helper is for distance models, not {other:?}"),
};
let device = test_device();
let param = |rows: &[Vec<f32>]| {
let cols = rows[0].len();
let flat: Vec<f32> = rows.iter().flat_map(|r| r.iter().copied()).collect();
Param::initialized(
ParamId::new(),
Tensor::<TestBackend, 2>::from_data(
burn::tensor::TensorData::new(flat, [rows.len(), cols]),
&device,
),
)
};
let kge = BurnKge::<TestBackend> {
entity: param(&ent_rows),
relation: param(&rel_rows),
};
let scores =
|t: Tensor<TestBackend, 2>| -> Vec<f32> { t.into_data().to_vec::<f32>().unwrap() };
let rank_by = |v: &[f32], higher_is_better: bool| -> Vec<usize> {
let mut idx: Vec<usize> = (0..v.len()).collect();
idx.sort_by(|&a, &b| {
let ord = if higher_is_better {
v[b].total_cmp(&v[a])
} else {
v[a].total_cmp(&v[b])
};
ord.then(a.cmp(&b))
});
idx
};
for r in 0..nr {
for x in 0..ne {
let ids = |i: usize| {
Tensor::<TestBackend, 1, Int>::from_data(
burn::tensor::TensorData::new(vec![i as i64], [1]),
&device,
)
};
let (rels, xs) = (ids(r), ids(x));
let burn_tails = scores(score_1n_kge(&kge, mt, dim, &xs, &rels));
let cpu_tails = cpu.score_all_tails(x, r);
assert_eq!(
rank_by(&burn_tails, true),
rank_by(&cpu_tails, false),
"{mt:?} tail ordering at rel={r} x={x}: burn={burn_tails:?} cpu={cpu_tails:?}"
);
let burn_heads = scores(score_1n_heads_kge(&kge, mt, dim, &rels, &xs));
let cpu_heads = cpu.score_all_heads(r, x);
assert_eq!(
rank_by(&burn_heads, true),
rank_by(&cpu_heads, false),
"{mt:?} head ordering at rel={r} x={x}: burn={burn_heads:?} cpu={cpu_heads:?}"
);
}
}
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_transe_ranks_match_cpu_reference() {
kge_ranks_match_cpu_reference(BurnModelType::TransE);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_rotate_ranks_match_cpu_reference() {
kge_ranks_match_cpu_reference(BurnModelType::RotatE);
}
#[cfg(feature = "burn-ndarray")]
fn kge_achieves_nonzero_mrr(mt: BurnModelType) {
let triples = vec![tid(0, 0, 1), tid(1, 0, 2), tid(2, 0, 3), tid(3, 0, 4)];
let config = BurnTrainConfig {
dim: 32,
epochs: 300,
batch_size: 4,
lr: 0.01,
..BurnTrainConfig::default()
};
let result = train_kge::<TestBackend>(&triples, 5, 1, mt, &config, &test_device());
let scorer = result.to_scorer();
let ds = crate::dataset::Dataset::new(
triples
.iter()
.map(|t| {
crate::dataset::Triple::new(
t.head.to_string(),
t.relation.to_string(),
t.tail.to_string(),
)
})
.collect(),
Vec::new(),
Vec::new(),
)
.into_interned();
let filter = crate::dataset::FilterIndex::from_dataset(&ds);
let metrics = crate::eval::evaluate_link_prediction(scorer.as_ref(), &triples, &filter, 5);
assert!(
metrics.mrr > 0.3,
"Burn {mt:?} should achieve MRR > 0.3, got {:.4}",
metrics.mrr
);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_transe_achieves_nonzero_mrr() {
kge_achieves_nonzero_mrr(BurnModelType::TransE);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_rotate_achieves_nonzero_mrr() {
kge_achieves_nonzero_mrr(BurnModelType::RotatE);
}
#[test]
#[cfg(feature = "burn-ndarray")]
fn burn_kge_to_scorer_builds_all_models() {
let triples = vec![tid(0, 0, 1), tid(1, 0, 2)];
let config = BurnTrainConfig {
dim: 8,
epochs: 5,
batch_size: 2,
..BurnTrainConfig::default()
};
for mt in [
BurnModelType::TransE,
BurnModelType::RotatE,
BurnModelType::ComplEx,
BurnModelType::DistMult,
] {
let result = train_kge::<TestBackend>(&triples, 3, 1, mt, &config, &test_device());
assert!(result.losses.iter().all(|l| l.is_finite()), "{mt:?}");
let scorer = result.to_scorer();
assert_eq!(scorer.num_entities(), 3, "{mt:?}");
}
}
}