Skip to main content

oxirs_embed/models/
transe.rs

1//! TransE: Translating Embeddings for Modeling Multi-relational Data
2//!
3//! TransE models relations as translations in the embedding space:
4//! h + r ≈ t for a true triple (h, r, t)
5//!
6//! Reference: Bordes et al. "Translating Embeddings for Modeling Multi-relational Data" (2013)
7
8use 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/// Serializable representation of a TransE model for persistence.
26#[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/// TransE embedding model
37#[derive(Debug, Clone)]
38pub struct TransE {
39    /// Base model functionality
40    base: BaseModel,
41    /// Entity embeddings matrix (num_entities × dimensions)
42    entity_embeddings: Array2<f64>,
43    /// Relation embeddings matrix (num_relations × dimensions)
44    relation_embeddings: Array2<f64>,
45    /// Whether embeddings have been initialized
46    embeddings_initialized: bool,
47    /// Distance metric for scoring
48    distance_metric: DistanceMetric,
49    /// Margin for ranking loss
50    margin: f64,
51}
52
53/// Distance metrics supported by TransE
54#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
55pub enum DistanceMetric {
56    /// L1 (Manhattan) distance
57    L1,
58    /// L2 (Euclidean) distance
59    L2,
60    /// Cosine distance (1 - cosine similarity)
61    Cosine,
62}
63
64impl TransE {
65    /// Create a new TransE model
66    pub fn new(config: ModelConfig) -> Self {
67        let base = BaseModel::new(config.clone());
68
69        // Get TransE-specific parameters
70        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, // Default to L2
75        };
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    /// Create a new TransE model with L1 (Manhattan) distance metric
90    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    /// Create a new TransE model with L2 (Euclidean) distance metric
98    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    /// Create a new TransE model with Cosine distance metric
106    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    /// Create a new TransE model with custom margin for ranking loss
114    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    /// Get the current distance metric
120    pub fn distance_metric(&self) -> DistanceMetric {
121        self.distance_metric
122    }
123
124    /// Get the current margin value
125    pub fn margin(&self) -> f64 {
126        self.margin
127    }
128
129    /// Initialize embeddings after entities and relations are known
130    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        // Initialize entity embeddings with Xavier initialization
146        self.entity_embeddings =
147            xavier_init((num_entities, dimensions), dimensions, dimensions, &mut rng);
148
149        // Initialize relation embeddings with Xavier initialization
150        self.relation_embeddings = xavier_init(
151            (num_relations, dimensions),
152            dimensions,
153            dimensions,
154            &mut rng,
155        );
156
157        // Normalize entity embeddings to unit sphere
158        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    /// Score a triple using TransE scoring function
168    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        // Compute h + r - t
183        let diff = &h + &r - t;
184
185        // Distance metric determines scoring (lower distance = higher score)
186        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                // For cosine distance, we compute 1 - cosine_similarity(h + r, t)
191                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 // Maximum distance for zero vectors
198                } else {
199                    let cosine_sim = dot_product / (norm_h_plus_r * norm_t);
200                    1.0 - cosine_sim.clamp(-1.0, 1.0) // Clamp to [-1, 1] and convert to distance
201                }
202            }
203        };
204
205        // Return negative distance as score (higher is better)
206        Ok(-distance)
207    }
208
209    /// Compute gradients for a training triple
210    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        // Get embeddings
222        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        // Compute differences
231        let pos_diff = &pos_h + &pos_r - pos_t;
232        let neg_diff = &neg_h + &neg_r - neg_t;
233
234        // Compute distances
235        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        // Check if we need to update (margin loss > 0)
262        let loss = self.margin + pos_distance - neg_distance;
263        if loss > 0.0 {
264            // Compute gradient direction based on distance metric
265            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            // Update gradients for positive triple (increase distance)
320            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            // Update gradients for negative triple (decrease distance)
327            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    /// Perform one training epoch
338    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        // Create shuffled batches
346        let mut shuffled_triples = self.base.triples.clone();
347        // Manual Fisher-Yates shuffle using scirs2-core
348        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                // Generate negative samples
360                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                    // Compute scores
366                    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                    // Convert scores to distances (negate because score = -distance)
372                    let pos_distance = -pos_score;
373                    let neg_distance = -neg_score;
374
375                    // Compute margin loss. margin_loss(positive_score, negative_score, margin)
376                    // = max(0, margin + negative_score - positive_score). For TransE the hinge
377                    // loss is max(0, margin + pos_distance - neg_distance), so pass neg_distance
378                    // as the positive-score slot and pos_distance as the negative-score slot to
379                    // match compute_gradients' internal `margin + pos_distance - neg_distance`.
380                    let loss = margin_loss(neg_distance, pos_distance, self.margin);
381                    batch_loss += loss;
382
383                    if loss > 0.0 {
384                        // Compute and accumulate gradients
385                        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            // Apply gradients
394            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 entity embeddings
410                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    /// Create a TransE model using [`ModelConfig::default`] (100-dimensional
422    /// embeddings, L2 distance, margin 1.0, learning rate 0.01).
423    ///
424    /// This mirrors the construction already used for "give me *a* TransE
425    /// model" call sites elsewhere in the crate (e.g. the model registry in
426    /// `persistence.rs`), and exists so `TransE` can satisfy generic bounds
427    /// like `M: EmbeddingModel + Default` used by quick smoke-test harnesses
428    /// such as [`crate::evaluation::kgc_evaluator::KgcEvaluationSuite::run_on_synthetic`].
429    /// Callers that care about a specific embedding dimension or learning
430    /// rate should use [`TransE::new`] with an explicit [`ModelConfig`]
431    /// instead.
432    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        // Initialize embeddings if needed
460        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            // Simple convergence check
480            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        // Add test triples
746        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        // Train
754        let stats = model.train(Some(5)).await?;
755        assert!(stats.epochs_completed > 0);
756
757        // Test embeddings
758        let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
759        assert_eq!(alice_emb.dimensions, 50);
760
761        // Test scoring
762        let score = model.score_triple(
763            "http://example.org/alice",
764            "http://example.org/knows",
765            "http://example.org/bob",
766        )?;
767
768        // Score should be a finite number
769        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        // Test L1 distance
782        let mut model_l1 = TransE::with_l1_distance(base_config.clone());
783        assert!(matches!(model_l1.distance_metric(), DistanceMetric::L1));
784
785        // Test L2 distance
786        let mut model_l2 = TransE::with_l2_distance(base_config.clone());
787        assert!(matches!(model_l2.distance_metric(), DistanceMetric::L2));
788
789        // Test Cosine distance
790        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        // Test custom margin
797        let model_margin = TransE::with_margin(base_config.clone(), 2.0);
798        assert_eq!(model_margin.margin(), 2.0);
799
800        // Add same triples to all models
801        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        // Train all models
811        model_l1.train(Some(3)).await?;
812        model_l2.train(Some(3)).await?;
813        model_cosine.train(Some(3)).await?;
814
815        // Test that all models produce finite scores
816        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        // Scores may differ due to different distance metrics
837        // This tests that the cosine distance implementation works
838        println!("L1 score: {score_l1}, L2 score: {score_l2}, Cosine score: {score_cosine}");
839
840        Ok(())
841    }
842
843    /// Regression: the TransE training loop must gate gradient updates on the
844    /// correct hinge loss `max(0, margin + pos_distance - neg_distance)`. The
845    /// previous code passed distances in the wrong argument order, inverting the
846    /// sign so that badly-violated triples (the ones most needing an update)
847    /// produced zero loss and were skipped.
848    #[test]
849    fn regression_transe_margin_loss_orientation() {
850        use crate::models::common::margin_loss;
851        let margin = 1.0;
852
853        // Badly-violated triple: positive distance far larger than negative.
854        // Correct hinge loss must be strictly positive so an update happens.
855        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        // Well-separated triple: positive close, negative far => zero loss.
865        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    /// Regression: save()/load() were no-ops that silently lost all trained
870    /// weights. A round-trip through disk must reproduce identical embeddings.
871    #[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        // Fresh, untrained model of the same type; load must restore weights.
893        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}