1use crate::models::serialization::{BaseModelSnapshot, MatrixF64};
9use crate::models::{common::*, BaseModel};
10use crate::{EmbeddingModel, ModelConfig, ModelStats, TrainingStats, Triple, Vector};
11use anyhow::{anyhow, Result};
12use async_trait::async_trait;
13use scirs2_core::ndarray_ext::{Array1, Array2};
14#[allow(unused_imports)]
15use scirs2_core::random::{Random, RngExt};
16use serde::{Deserialize, Serialize};
17use std::fs::File;
18use std::io::{BufReader, BufWriter};
19use std::ops::{AddAssign, SubAssign};
20use std::path::Path;
21use std::time::Instant;
22use tracing::{debug, info};
23use uuid::Uuid;
24
25#[derive(Debug, Serialize, Deserialize)]
27struct TransESerializable {
28 base: BaseModelSnapshot,
29 entity_embeddings: MatrixF64,
30 relation_embeddings: MatrixF64,
31 embeddings_initialized: bool,
32 distance_metric: DistanceMetric,
33 margin: f64,
34}
35
36#[derive(Debug, Clone)]
38pub struct TransE {
39 base: BaseModel,
41 entity_embeddings: Array2<f64>,
43 relation_embeddings: Array2<f64>,
45 embeddings_initialized: bool,
47 distance_metric: DistanceMetric,
49 margin: f64,
51}
52
53#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
55pub enum DistanceMetric {
56 L1,
58 L2,
60 Cosine,
62}
63
64impl TransE {
65 pub fn new(config: ModelConfig) -> Self {
67 let base = BaseModel::new(config.clone());
68
69 let distance_metric = match config.model_params.get("distance_metric") {
71 Some(0.0) => DistanceMetric::L1,
72 Some(1.0) => DistanceMetric::L2,
73 Some(2.0) => DistanceMetric::Cosine,
74 _ => DistanceMetric::L2, };
76
77 let margin = config.model_params.get("margin").copied().unwrap_or(1.0);
78
79 Self {
80 base,
81 entity_embeddings: Array2::zeros((0, config.dimensions)),
82 relation_embeddings: Array2::zeros((0, config.dimensions)),
83 embeddings_initialized: false,
84 distance_metric,
85 margin,
86 }
87 }
88
89 pub fn with_l1_distance(mut config: ModelConfig) -> Self {
91 config
92 .model_params
93 .insert("distance_metric".to_string(), 0.0);
94 Self::new(config)
95 }
96
97 pub fn with_l2_distance(mut config: ModelConfig) -> Self {
99 config
100 .model_params
101 .insert("distance_metric".to_string(), 1.0);
102 Self::new(config)
103 }
104
105 pub fn with_cosine_distance(mut config: ModelConfig) -> Self {
107 config
108 .model_params
109 .insert("distance_metric".to_string(), 2.0);
110 Self::new(config)
111 }
112
113 pub fn with_margin(mut config: ModelConfig, margin: f64) -> Self {
115 config.model_params.insert("margin".to_string(), margin);
116 Self::new(config)
117 }
118
119 pub fn distance_metric(&self) -> DistanceMetric {
121 self.distance_metric
122 }
123
124 pub fn margin(&self) -> f64 {
126 self.margin
127 }
128
129 fn initialize_embeddings(&mut self) {
131 if self.embeddings_initialized {
132 return;
133 }
134
135 let num_entities = self.base.num_entities();
136 let num_relations = self.base.num_relations();
137 let dimensions = self.base.config.dimensions;
138
139 if num_entities == 0 || num_relations == 0 {
140 return;
141 }
142
143 let mut rng = Random::default();
144
145 self.entity_embeddings =
147 xavier_init((num_entities, dimensions), dimensions, dimensions, &mut rng);
148
149 self.relation_embeddings = xavier_init(
151 (num_relations, dimensions),
152 dimensions,
153 dimensions,
154 &mut rng,
155 );
156
157 normalize_embeddings(&mut self.entity_embeddings);
159
160 self.embeddings_initialized = true;
161 debug!(
162 "Initialized TransE embeddings: {} entities, {} relations, {} dimensions",
163 num_entities, num_relations, dimensions
164 );
165 }
166
167 fn score_triple_ids(
169 &self,
170 subject_id: usize,
171 predicate_id: usize,
172 object_id: usize,
173 ) -> Result<f64> {
174 if !self.embeddings_initialized {
175 return Err(anyhow!("Model not trained"));
176 }
177
178 let h = self.entity_embeddings.row(subject_id);
179 let r = self.relation_embeddings.row(predicate_id);
180 let t = self.entity_embeddings.row(object_id);
181
182 let diff = &h + &r - t;
184
185 let distance = match self.distance_metric {
187 DistanceMetric::L1 => diff.mapv(|x| x.abs()).sum(),
188 DistanceMetric::L2 => diff.mapv(|x| x * x).sum().sqrt(),
189 DistanceMetric::Cosine => {
190 let h_plus_r = &h + &r;
192 let dot_product = (&h_plus_r * &t).sum();
193 let norm_h_plus_r = h_plus_r.mapv(|x| x * x).sum().sqrt();
194 let norm_t = t.mapv(|x| x * x).sum().sqrt();
195
196 if norm_h_plus_r == 0.0 || norm_t == 0.0 {
197 1.0 } else {
199 let cosine_sim = dot_product / (norm_h_plus_r * norm_t);
200 1.0 - cosine_sim.clamp(-1.0, 1.0) }
202 }
203 };
204
205 Ok(-distance)
207 }
208
209 fn compute_gradients(
211 &self,
212 pos_triple: (usize, usize, usize),
213 neg_triple: (usize, usize, usize),
214 ) -> Result<(Array2<f64>, Array2<f64>)> {
215 let (pos_s, pos_p, pos_o) = pos_triple;
216 let (neg_s, neg_p, neg_o) = neg_triple;
217
218 let mut entity_grads = Array2::zeros(self.entity_embeddings.raw_dim());
219 let mut relation_grads = Array2::zeros(self.relation_embeddings.raw_dim());
220
221 let pos_h = self.entity_embeddings.row(pos_s);
223 let pos_r = self.relation_embeddings.row(pos_p);
224 let pos_t = self.entity_embeddings.row(pos_o);
225
226 let neg_h = self.entity_embeddings.row(neg_s);
227 let neg_r = self.relation_embeddings.row(neg_p);
228 let neg_t = self.entity_embeddings.row(neg_o);
229
230 let pos_diff = &pos_h + &pos_r - pos_t;
232 let neg_diff = &neg_h + &neg_r - neg_t;
233
234 let pos_distance = match self.distance_metric {
236 DistanceMetric::L1 => pos_diff.mapv(|x| x.abs()).sum(),
237 DistanceMetric::L2 => pos_diff.mapv(|x| x * x).sum().sqrt(),
238 DistanceMetric::Cosine => {
239 let norm = pos_diff.mapv(|x| x * x).sum().sqrt();
240 if norm > 1e-10 {
241 1.0 - (pos_diff.dot(&pos_diff) / (norm * norm)).clamp(-1.0, 1.0)
242 } else {
243 0.0
244 }
245 }
246 };
247
248 let neg_distance = match self.distance_metric {
249 DistanceMetric::L1 => neg_diff.mapv(|x| x.abs()).sum(),
250 DistanceMetric::L2 => neg_diff.mapv(|x| x * x).sum().sqrt(),
251 DistanceMetric::Cosine => {
252 let norm = neg_diff.mapv(|x| x * x).sum().sqrt();
253 if norm > 1e-10 {
254 1.0 - (neg_diff.dot(&neg_diff) / (norm * norm)).clamp(-1.0, 1.0)
255 } else {
256 0.0
257 }
258 }
259 };
260
261 let loss = self.margin + pos_distance - neg_distance;
263 if loss > 0.0 {
264 let pos_grad_direction = match self.distance_metric {
266 DistanceMetric::L1 => pos_diff.mapv(|x| {
267 if x > 0.0 {
268 1.0
269 } else if x < 0.0 {
270 -1.0
271 } else {
272 0.0
273 }
274 }),
275 DistanceMetric::L2 => {
276 if pos_distance > 1e-10 {
277 &pos_diff / pos_distance
278 } else {
279 Array1::zeros(pos_diff.len())
280 }
281 }
282 DistanceMetric::Cosine => {
283 let norm_sq = pos_diff.mapv(|x| x * x).sum();
284 if norm_sq > 1e-10 {
285 &pos_diff / norm_sq.sqrt()
286 } else {
287 Array1::zeros(pos_diff.len())
288 }
289 }
290 };
291
292 let neg_grad_direction = match self.distance_metric {
293 DistanceMetric::L1 => neg_diff.mapv(|x| {
294 if x > 0.0 {
295 1.0
296 } else if x < 0.0 {
297 -1.0
298 } else {
299 0.0
300 }
301 }),
302 DistanceMetric::L2 => {
303 if neg_distance > 1e-10 {
304 &neg_diff / neg_distance
305 } else {
306 Array1::zeros(neg_diff.len())
307 }
308 }
309 DistanceMetric::Cosine => {
310 let norm_sq = neg_diff.mapv(|x| x * x).sum();
311 if norm_sq > 1e-10 {
312 &neg_diff / norm_sq.sqrt()
313 } else {
314 Array1::zeros(neg_diff.len())
315 }
316 }
317 };
318
319 entity_grads.row_mut(pos_s).add_assign(&pos_grad_direction);
321 relation_grads
322 .row_mut(pos_p)
323 .add_assign(&pos_grad_direction);
324 entity_grads.row_mut(pos_o).sub_assign(&pos_grad_direction);
325
326 entity_grads.row_mut(neg_s).sub_assign(&neg_grad_direction);
328 relation_grads
329 .row_mut(neg_p)
330 .sub_assign(&neg_grad_direction);
331 entity_grads.row_mut(neg_o).add_assign(&neg_grad_direction);
332 }
333
334 Ok((entity_grads, relation_grads))
335 }
336
337 async fn train_epoch(&mut self, learning_rate: f64) -> Result<f64> {
339 let mut rng = Random::default();
340
341 let mut total_loss = 0.0;
342 let num_batches = (self.base.triples.len() + self.base.config.batch_size - 1)
343 / self.base.config.batch_size;
344
345 let mut shuffled_triples = self.base.triples.clone();
347 for i in (1..shuffled_triples.len()).rev() {
349 let j = rng.random_range(0..i + 1);
350 shuffled_triples.swap(i, j);
351 }
352
353 for batch_triples in shuffled_triples.chunks(self.base.config.batch_size) {
354 let mut batch_entity_grads = Array2::zeros(self.entity_embeddings.raw_dim());
355 let mut batch_relation_grads = Array2::zeros(self.relation_embeddings.raw_dim());
356 let mut batch_loss = 0.0;
357
358 for &pos_triple in batch_triples {
359 let neg_samples = self
361 .base
362 .generate_negative_samples(self.base.config.negative_samples, &mut rng);
363
364 for neg_triple in neg_samples {
365 let pos_score =
367 self.score_triple_ids(pos_triple.0, pos_triple.1, pos_triple.2)?;
368 let neg_score =
369 self.score_triple_ids(neg_triple.0, neg_triple.1, neg_triple.2)?;
370
371 let pos_distance = -pos_score;
373 let neg_distance = -neg_score;
374
375 let loss = margin_loss(neg_distance, pos_distance, self.margin);
381 batch_loss += loss;
382
383 if loss > 0.0 {
384 let (entity_grads, relation_grads) =
386 self.compute_gradients(pos_triple, neg_triple)?;
387 batch_entity_grads += &entity_grads;
388 batch_relation_grads += &relation_grads;
389 }
390 }
391 }
392
393 if batch_loss > 0.0 {
395 gradient_update(
396 &mut self.entity_embeddings,
397 &batch_entity_grads,
398 learning_rate,
399 self.base.config.l2_reg,
400 );
401
402 gradient_update(
403 &mut self.relation_embeddings,
404 &batch_relation_grads,
405 learning_rate,
406 self.base.config.l2_reg,
407 );
408
409 normalize_embeddings(&mut self.entity_embeddings);
411 }
412
413 total_loss += batch_loss;
414 }
415
416 Ok(total_loss / num_batches as f64)
417 }
418}
419
420impl Default for TransE {
421 fn default() -> Self {
433 Self::new(ModelConfig::default())
434 }
435}
436
437#[async_trait]
438impl EmbeddingModel for TransE {
439 fn config(&self) -> &ModelConfig {
440 &self.base.config
441 }
442
443 fn model_id(&self) -> &Uuid {
444 &self.base.model_id
445 }
446
447 fn model_type(&self) -> &'static str {
448 "TransE"
449 }
450
451 fn add_triple(&mut self, triple: Triple) -> Result<()> {
452 self.base.add_triple(triple)
453 }
454
455 async fn train(&mut self, epochs: Option<usize>) -> Result<TrainingStats> {
456 let start_time = Instant::now();
457 let max_epochs = epochs.unwrap_or(self.base.config.max_epochs);
458
459 self.initialize_embeddings();
461
462 if !self.embeddings_initialized {
463 return Err(anyhow!("No training data available"));
464 }
465
466 let mut loss_history = Vec::new();
467 let learning_rate = self.base.config.learning_rate;
468
469 info!("Starting TransE training for {} epochs", max_epochs);
470
471 for epoch in 0..max_epochs {
472 let epoch_loss = self.train_epoch(learning_rate).await?;
473 loss_history.push(epoch_loss);
474
475 if epoch % 100 == 0 {
476 debug!("Epoch {}: loss = {:.6}", epoch, epoch_loss);
477 }
478
479 if epoch > 10 && epoch_loss < 1e-6 {
481 info!("Converged at epoch {} with loss {:.6}", epoch, epoch_loss);
482 break;
483 }
484 }
485
486 self.base.mark_trained();
487 let training_time = start_time.elapsed().as_secs_f64();
488
489 Ok(TrainingStats {
490 epochs_completed: loss_history.len(),
491 final_loss: loss_history.last().copied().unwrap_or(0.0),
492 training_time_seconds: training_time,
493 convergence_achieved: loss_history.last().copied().unwrap_or(f64::INFINITY) < 1e-6,
494 loss_history,
495 })
496 }
497
498 fn get_entity_embedding(&self, entity: &str) -> Result<Vector> {
499 if !self.embeddings_initialized {
500 return Err(anyhow!("Model not trained"));
501 }
502
503 let entity_id = self
504 .base
505 .get_entity_id(entity)
506 .ok_or_else(|| anyhow!("Entity not found: {}", entity))?;
507
508 let embedding = self.entity_embeddings.row(entity_id).to_owned();
509 Ok(ndarray_to_vector(&embedding))
510 }
511
512 fn get_relation_embedding(&self, relation: &str) -> Result<Vector> {
513 if !self.embeddings_initialized {
514 return Err(anyhow!("Model not trained"));
515 }
516
517 let relation_id = self
518 .base
519 .get_relation_id(relation)
520 .ok_or_else(|| anyhow!("Relation not found: {}", relation))?;
521
522 let embedding = self.relation_embeddings.row(relation_id).to_owned();
523 Ok(ndarray_to_vector(&embedding))
524 }
525
526 fn score_triple(&self, subject: &str, predicate: &str, object: &str) -> Result<f64> {
527 let subject_id = self
528 .base
529 .get_entity_id(subject)
530 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
531 let predicate_id = self
532 .base
533 .get_relation_id(predicate)
534 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
535 let object_id = self
536 .base
537 .get_entity_id(object)
538 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
539
540 self.score_triple_ids(subject_id, predicate_id, object_id)
541 }
542
543 fn predict_objects(
544 &self,
545 subject: &str,
546 predicate: &str,
547 k: usize,
548 ) -> Result<Vec<(String, f64)>> {
549 if !self.embeddings_initialized {
550 return Err(anyhow!("Model not trained"));
551 }
552
553 let subject_id = self
554 .base
555 .get_entity_id(subject)
556 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
557 let predicate_id = self
558 .base
559 .get_relation_id(predicate)
560 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
561
562 let mut scores = Vec::new();
563
564 for object_id in 0..self.base.num_entities() {
565 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
566 let object_name = self
567 .base
568 .get_entity(object_id)
569 .expect("entity should exist for valid id")
570 .clone();
571 scores.push((object_name, score));
572 }
573
574 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
575 scores.truncate(k);
576
577 Ok(scores)
578 }
579
580 fn predict_subjects(
581 &self,
582 predicate: &str,
583 object: &str,
584 k: usize,
585 ) -> Result<Vec<(String, f64)>> {
586 if !self.embeddings_initialized {
587 return Err(anyhow!("Model not trained"));
588 }
589
590 let predicate_id = self
591 .base
592 .get_relation_id(predicate)
593 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
594 let object_id = self
595 .base
596 .get_entity_id(object)
597 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
598
599 let mut scores = Vec::new();
600
601 for subject_id in 0..self.base.num_entities() {
602 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
603 let subject_name = self
604 .base
605 .get_entity(subject_id)
606 .expect("entity should exist for valid id")
607 .clone();
608 scores.push((subject_name, score));
609 }
610
611 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
612 scores.truncate(k);
613
614 Ok(scores)
615 }
616
617 fn predict_relations(
618 &self,
619 subject: &str,
620 object: &str,
621 k: usize,
622 ) -> Result<Vec<(String, f64)>> {
623 if !self.embeddings_initialized {
624 return Err(anyhow!("Model not trained"));
625 }
626
627 let subject_id = self
628 .base
629 .get_entity_id(subject)
630 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
631 let object_id = self
632 .base
633 .get_entity_id(object)
634 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
635
636 let mut scores = Vec::new();
637
638 for predicate_id in 0..self.base.num_relations() {
639 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
640 let predicate_name = self
641 .base
642 .get_relation(predicate_id)
643 .expect("relation should exist for valid id")
644 .clone();
645 scores.push((predicate_name, score));
646 }
647
648 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
649 scores.truncate(k);
650
651 Ok(scores)
652 }
653
654 fn get_entities(&self) -> Vec<String> {
655 self.base.get_entities()
656 }
657
658 fn get_relations(&self) -> Vec<String> {
659 self.base.get_relations()
660 }
661
662 fn get_stats(&self) -> ModelStats {
663 self.base.get_stats("TransE")
664 }
665
666 fn save(&self, path: &str) -> Result<()> {
667 info!("Saving TransE model to {}", path);
668
669 let serializable = TransESerializable {
670 base: BaseModelSnapshot::capture(&self.base),
671 entity_embeddings: MatrixF64::from_array(&self.entity_embeddings),
672 relation_embeddings: MatrixF64::from_array(&self.relation_embeddings),
673 embeddings_initialized: self.embeddings_initialized,
674 distance_metric: self.distance_metric,
675 margin: self.margin,
676 };
677
678 let file = File::create(path)
679 .map_err(|e| anyhow!("Failed to create model file {}: {}", path, e))?;
680 let writer = BufWriter::new(file);
681 oxicode::serde::encode_into_std_write(&serializable, writer, oxicode::config::standard())
682 .map_err(|e| anyhow!("Failed to serialize TransE model: {}", e))?;
683
684 info!("TransE model saved successfully");
685 Ok(())
686 }
687
688 fn load(&mut self, path: &str) -> Result<()> {
689 info!("Loading TransE model from {}", path);
690
691 if !Path::new(path).exists() {
692 return Err(anyhow!("Model file not found: {}", path));
693 }
694
695 let file =
696 File::open(path).map_err(|e| anyhow!("Failed to open model file {}: {}", path, e))?;
697 let reader = BufReader::new(file);
698 let (serializable, _): (TransESerializable, _) =
699 oxicode::serde::decode_from_std_read(reader, oxicode::config::standard())
700 .map_err(|e| anyhow!("Failed to deserialize TransE model: {}", e))?;
701
702 self.entity_embeddings = serializable.entity_embeddings.to_array()?;
703 self.relation_embeddings = serializable.relation_embeddings.to_array()?;
704 self.embeddings_initialized = serializable.embeddings_initialized;
705 self.distance_metric = serializable.distance_metric;
706 self.margin = serializable.margin;
707 serializable.base.restore_into(&mut self.base);
708
709 info!("TransE model loaded successfully");
710 Ok(())
711 }
712
713 fn clear(&mut self) {
714 self.base.clear();
715 self.entity_embeddings = Array2::zeros((0, self.base.config.dimensions));
716 self.relation_embeddings = Array2::zeros((0, self.base.config.dimensions));
717 self.embeddings_initialized = false;
718 }
719
720 fn is_trained(&self) -> bool {
721 self.base.is_trained
722 }
723
724 async fn encode(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>> {
725 Err(anyhow!(
726 "TransE is a knowledge graph embedding model and does not support text encoding"
727 ))
728 }
729}
730
731#[cfg(test)]
732mod tests {
733 use super::*;
734 use crate::NamedNode;
735
736 #[tokio::test]
737 async fn test_transe_basic() -> Result<()> {
738 let config = ModelConfig::default()
739 .with_dimensions(50)
740 .with_max_epochs(10)
741 .with_seed(42);
742
743 let mut model = TransE::new(config);
744
745 let alice = NamedNode::new("http://example.org/alice")?;
747 let knows = NamedNode::new("http://example.org/knows")?;
748 let bob = NamedNode::new("http://example.org/bob")?;
749
750 model.add_triple(Triple::new(alice.clone(), knows.clone(), bob.clone()))?;
751 model.add_triple(Triple::new(bob.clone(), knows.clone(), alice.clone()))?;
752
753 let stats = model.train(Some(5)).await?;
755 assert!(stats.epochs_completed > 0);
756
757 let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
759 assert_eq!(alice_emb.dimensions, 50);
760
761 let score = model.score_triple(
763 "http://example.org/alice",
764 "http://example.org/knows",
765 "http://example.org/bob",
766 )?;
767
768 assert!(score.is_finite());
770
771 Ok(())
772 }
773
774 #[tokio::test]
775 async fn test_transe_distance_metrics() -> Result<()> {
776 let base_config = ModelConfig::default()
777 .with_dimensions(10)
778 .with_max_epochs(5)
779 .with_seed(42);
780
781 let mut model_l1 = TransE::with_l1_distance(base_config.clone());
783 assert!(matches!(model_l1.distance_metric(), DistanceMetric::L1));
784
785 let mut model_l2 = TransE::with_l2_distance(base_config.clone());
787 assert!(matches!(model_l2.distance_metric(), DistanceMetric::L2));
788
789 let mut model_cosine = TransE::with_cosine_distance(base_config.clone());
791 assert!(matches!(
792 model_cosine.distance_metric(),
793 DistanceMetric::Cosine
794 ));
795
796 let model_margin = TransE::with_margin(base_config.clone(), 2.0);
798 assert_eq!(model_margin.margin(), 2.0);
799
800 let alice = NamedNode::new("http://example.org/alice")?;
802 let knows = NamedNode::new("http://example.org/knows")?;
803 let bob = NamedNode::new("http://example.org/bob")?;
804 let triple = Triple::new(alice, knows, bob);
805
806 model_l1.add_triple(triple.clone())?;
807 model_l2.add_triple(triple.clone())?;
808 model_cosine.add_triple(triple.clone())?;
809
810 model_l1.train(Some(3)).await?;
812 model_l2.train(Some(3)).await?;
813 model_cosine.train(Some(3)).await?;
814
815 let score_l1 = model_l1.score_triple(
817 "http://example.org/alice",
818 "http://example.org/knows",
819 "http://example.org/bob",
820 )?;
821 let score_l2 = model_l2.score_triple(
822 "http://example.org/alice",
823 "http://example.org/knows",
824 "http://example.org/bob",
825 )?;
826 let score_cosine = model_cosine.score_triple(
827 "http://example.org/alice",
828 "http://example.org/knows",
829 "http://example.org/bob",
830 )?;
831
832 assert!(score_l1.is_finite());
833 assert!(score_l2.is_finite());
834 assert!(score_cosine.is_finite());
835
836 println!("L1 score: {score_l1}, L2 score: {score_l2}, Cosine score: {score_cosine}");
839
840 Ok(())
841 }
842
843 #[test]
849 fn regression_transe_margin_loss_orientation() {
850 use crate::models::common::margin_loss;
851 let margin = 1.0;
852
853 let pos_distance = 5.0;
856 let neg_distance = 1.0;
857 let loss = margin_loss(neg_distance, pos_distance, margin);
858 assert!(
859 loss > 0.0,
860 "violated triple must yield positive hinge loss, got {loss}"
861 );
862 assert!((loss - (margin + pos_distance - neg_distance)).abs() < 1e-9);
863
864 let good = margin_loss(5.0, 0.0, margin);
866 assert_eq!(good, 0.0, "well-separated triple must yield zero loss");
867 }
868
869 #[tokio::test]
872 async fn regression_transe_save_load_roundtrip() -> Result<()> {
873 let config = ModelConfig::default()
874 .with_dimensions(16)
875 .with_max_epochs(5)
876 .with_seed(7);
877 let mut model = TransE::new(config);
878
879 let alice = NamedNode::new("http://example.org/alice")?;
880 let knows = NamedNode::new("http://example.org/knows")?;
881 let bob = NamedNode::new("http://example.org/bob")?;
882 model.add_triple(Triple::new(alice.clone(), knows.clone(), bob.clone()))?;
883 model.add_triple(Triple::new(bob.clone(), knows.clone(), alice.clone()))?;
884 model.train(Some(5)).await?;
885
886 let before = model.get_entity_embedding("http://example.org/alice")?;
887
888 let path = std::env::temp_dir().join(format!("transe-roundtrip-{}.bin", Uuid::new_v4()));
889 let path_str = path.to_string_lossy().to_string();
890 model.save(&path_str)?;
891
892 let mut restored = TransE::new(ModelConfig::default());
894 restored.load(&path_str)?;
895
896 assert!(restored.is_trained());
897 let after = restored.get_entity_embedding("http://example.org/alice")?;
898 assert_eq!(before.dimensions, after.dimensions);
899 for (x, y) in before.values.iter().zip(after.values.iter()) {
900 assert!(
901 (x - y).abs() < 1e-9,
902 "embedding mismatch after load: {x} vs {y}"
903 );
904 }
905
906 let _ = std::fs::remove_file(&path);
907 Ok(())
908 }
909}