Skip to main content

oxirs_embed/models/
distmult.rs

1//! DistMult: Embedding Entities and Relations for Learning and Inference in Knowledge Bases
2//!
3//! DistMult is a bilinear model that uses element-wise multiplication:
4//! score(h, r, t) = h^T diag(r) t = sum(h * r * t)
5//!
6//! Note: DistMult can only model symmetric relations due to its formulation.
7//!
8//! Reference: Yang et al. "Embedding Entities and Relations for Learning and Inference in Knowledge Bases" (2014)
9
10use crate::models::serialization::{BaseModelSnapshot, MatrixF64};
11use crate::models::{common::*, BaseModel};
12use crate::{EmbeddingModel, ModelConfig, ModelStats, TrainingStats, Triple, Vector};
13use anyhow::{anyhow, Result};
14use async_trait::async_trait;
15use scirs2_core::ndarray_ext::{Array1, Array2};
16#[allow(unused_imports)]
17use scirs2_core::random::{Random, RngExt};
18use serde::{Deserialize, Serialize};
19use std::fs::File;
20use std::io::{BufReader, BufWriter};
21use std::ops::AddAssign;
22use std::path::Path;
23use std::time::Instant;
24use tracing::{debug, info, warn};
25
26/// Serializable representation of a DistMult model for persistence.
27#[derive(Debug, Serialize, Deserialize)]
28struct DistMultSerializable {
29    base: BaseModelSnapshot,
30    entity_embeddings: MatrixF64,
31    relation_embeddings: MatrixF64,
32    embeddings_initialized: bool,
33    dropout_rate: f64,
34    loss_function: LossFunction,
35}
36use uuid::Uuid;
37
38/// DistMult embedding model
39#[derive(Debug)]
40pub struct DistMult {
41    /// Base model functionality
42    base: BaseModel,
43    /// Entity embeddings matrix (num_entities × dimensions)
44    entity_embeddings: Array2<f64>,
45    /// Relation embeddings matrix (num_relations × dimensions)
46    relation_embeddings: Array2<f64>,
47    /// Whether embeddings have been initialized
48    embeddings_initialized: bool,
49    /// Whether to apply dropout during training
50    #[allow(dead_code)]
51    dropout_rate: f64,
52    /// Loss function type
53    loss_function: LossFunction,
54}
55
56/// Loss function types for DistMult
57#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
58pub enum LossFunction {
59    /// Logistic loss (binary cross-entropy)
60    Logistic,
61    /// Margin-based ranking loss
62    MarginRanking,
63    /// Squared loss
64    SquaredLoss,
65}
66
67impl DistMult {
68    /// Create a new DistMult model
69    pub fn new(config: ModelConfig) -> Self {
70        let base = BaseModel::new(config.clone());
71
72        // Get DistMult-specific parameters
73        let dropout_rate = config
74            .model_params
75            .get("dropout_rate")
76            .copied()
77            .unwrap_or(0.0);
78
79        let loss_function = match config.model_params.get("loss_function") {
80            Some(0.0) => LossFunction::Logistic,
81            Some(1.0) => LossFunction::MarginRanking,
82            Some(2.0) => LossFunction::SquaredLoss,
83            _ => LossFunction::Logistic, // Default to logistic
84        };
85
86        Self {
87            base,
88            entity_embeddings: Array2::zeros((0, config.dimensions)),
89            relation_embeddings: Array2::zeros((0, config.dimensions)),
90            embeddings_initialized: false,
91            dropout_rate,
92            loss_function,
93        }
94    }
95
96    /// Initialize embeddings
97    fn initialize_embeddings(&mut self) {
98        if self.embeddings_initialized {
99            return;
100        }
101
102        let num_entities = self.base.num_entities();
103        let num_relations = self.base.num_relations();
104        let dimensions = self.base.config.dimensions;
105
106        if num_entities == 0 || num_relations == 0 {
107            return;
108        }
109
110        let mut rng = Random::default();
111
112        // Initialize embeddings with Xavier initialization
113        self.entity_embeddings =
114            xavier_init((num_entities, dimensions), dimensions, dimensions, &mut rng);
115
116        self.relation_embeddings = xavier_init(
117            (num_relations, dimensions),
118            dimensions,
119            dimensions,
120            &mut rng,
121        );
122
123        // For DistMult, it's common to normalize entity embeddings initially
124        normalize_embeddings(&mut self.entity_embeddings);
125
126        self.embeddings_initialized = true;
127        debug!(
128            "Initialized DistMult embeddings: {} entities, {} relations, {} dimensions",
129            num_entities, num_relations, dimensions
130        );
131    }
132
133    /// Score a triple using DistMult scoring function: sum(h * r * t)
134    fn score_triple_ids(
135        &self,
136        subject_id: usize,
137        predicate_id: usize,
138        object_id: usize,
139    ) -> Result<f64> {
140        if !self.embeddings_initialized {
141            return Err(anyhow!("Model not trained"));
142        }
143
144        let h = self.entity_embeddings.row(subject_id);
145        let r = self.relation_embeddings.row(predicate_id);
146        let t = self.entity_embeddings.row(object_id);
147
148        // DistMult score: sum of element-wise multiplication
149        let score = (&h * &r * t).sum();
150
151        Ok(score)
152    }
153
154    /// Apply dropout to embeddings during training
155    #[allow(dead_code)]
156    fn apply_dropout(&self, embeddings: &Array1<f64>, rng: &mut Random) -> Array1<f64> {
157        if self.dropout_rate > 0.0 {
158            embeddings.mapv(|x| {
159                if rng.random_f64() < self.dropout_rate {
160                    0.0
161                } else {
162                    x / (1.0 - self.dropout_rate)
163                }
164            })
165        } else {
166            embeddings.to_owned()
167        }
168    }
169
170    /// Compute gradients for DistMult model
171    fn compute_gradients(
172        &self,
173        pos_triple: (usize, usize, usize),
174        neg_triple: (usize, usize, usize),
175        pos_score: f64,
176        neg_score: f64,
177    ) -> Result<(Array2<f64>, Array2<f64>)> {
178        let mut entity_grads = Array2::zeros(self.entity_embeddings.raw_dim());
179        let mut relation_grads = Array2::zeros(self.relation_embeddings.raw_dim());
180
181        match self.loss_function {
182            LossFunction::Logistic => {
183                // Logistic loss gradients
184                let pos_sigmoid = sigmoid(pos_score);
185                let neg_sigmoid = sigmoid(neg_score);
186
187                let pos_grad_coeff = pos_sigmoid - 1.0; // d/dx log(sigmoid(x)) = sigmoid(x) - 1
188                let neg_grad_coeff = neg_sigmoid; // d/dx log(1 - sigmoid(x)) = sigmoid(x)
189
190                self.add_triple_gradients(
191                    pos_triple,
192                    pos_grad_coeff,
193                    &mut entity_grads,
194                    &mut relation_grads,
195                );
196                self.add_triple_gradients(
197                    neg_triple,
198                    neg_grad_coeff,
199                    &mut entity_grads,
200                    &mut relation_grads,
201                );
202            }
203            LossFunction::MarginRanking => {
204                // Margin ranking loss: max(0, margin + neg_score - pos_score)
205                let margin = self
206                    .base
207                    .config
208                    .model_params
209                    .get("margin")
210                    .copied()
211                    .unwrap_or(1.0);
212                let loss = margin + neg_score - pos_score;
213
214                if loss > 0.0 {
215                    // Gradients: -1 for positive triple, +1 for negative triple
216                    self.add_triple_gradients(
217                        pos_triple,
218                        -1.0,
219                        &mut entity_grads,
220                        &mut relation_grads,
221                    );
222                    self.add_triple_gradients(
223                        neg_triple,
224                        1.0,
225                        &mut entity_grads,
226                        &mut relation_grads,
227                    );
228                }
229            }
230            LossFunction::SquaredLoss => {
231                // Squared loss: (1 - pos_score)^2 + (0 - neg_score)^2
232                let pos_grad_coeff = -2.0 * (1.0 - pos_score);
233                let neg_grad_coeff = -2.0 * neg_score;
234
235                self.add_triple_gradients(
236                    pos_triple,
237                    pos_grad_coeff,
238                    &mut entity_grads,
239                    &mut relation_grads,
240                );
241                self.add_triple_gradients(
242                    neg_triple,
243                    neg_grad_coeff,
244                    &mut entity_grads,
245                    &mut relation_grads,
246                );
247            }
248        }
249
250        Ok((entity_grads, relation_grads))
251    }
252
253    /// Add gradients for a single triple
254    fn add_triple_gradients(
255        &self,
256        triple: (usize, usize, usize),
257        grad_coeff: f64,
258        entity_grads: &mut Array2<f64>,
259        relation_grads: &mut Array2<f64>,
260    ) {
261        let (s, p, o) = triple;
262
263        let h = self.entity_embeddings.row(s);
264        let r = self.relation_embeddings.row(p);
265        let t = self.entity_embeddings.row(o);
266
267        // DistMult gradients:
268        // ∂score/∂h = r * t
269        // ∂score/∂r = h * t
270        // ∂score/∂t = h * r
271
272        let h_grad = (&r * &t) * grad_coeff;
273        let r_grad = (&h * &t) * grad_coeff;
274        let t_grad = (&h * &r) * grad_coeff;
275
276        entity_grads.row_mut(s).add_assign(&h_grad);
277        relation_grads.row_mut(p).add_assign(&r_grad);
278        entity_grads.row_mut(o).add_assign(&t_grad);
279    }
280
281    /// Perform one training epoch
282    async fn train_epoch(&mut self, learning_rate: f64) -> Result<f64> {
283        let mut rng = Random::default();
284
285        let mut total_loss = 0.0;
286        let num_batches = (self.base.triples.len() + self.base.config.batch_size - 1)
287            / self.base.config.batch_size;
288
289        // Create shuffled batches
290        let mut shuffled_triples = self.base.triples.clone();
291        // Manual Fisher-Yates shuffle using scirs2-core
292        for i in (1..shuffled_triples.len()).rev() {
293            let j = rng.random_range(0..i + 1);
294            shuffled_triples.swap(i, j);
295        }
296
297        for batch_triples in shuffled_triples.chunks(self.base.config.batch_size) {
298            let mut batch_entity_grads = Array2::zeros(self.entity_embeddings.raw_dim());
299            let mut batch_relation_grads = Array2::zeros(self.relation_embeddings.raw_dim());
300            let mut batch_loss = 0.0;
301
302            for &pos_triple in batch_triples {
303                // Generate negative samples
304                let neg_samples = self
305                    .base
306                    .generate_negative_samples(self.base.config.negative_samples, &mut rng);
307
308                for neg_triple in neg_samples {
309                    // Compute scores
310                    let pos_score =
311                        self.score_triple_ids(pos_triple.0, pos_triple.1, pos_triple.2)?;
312                    let neg_score =
313                        self.score_triple_ids(neg_triple.0, neg_triple.1, neg_triple.2)?;
314
315                    // Compute loss
316                    let triple_loss = match self.loss_function {
317                        LossFunction::Logistic => {
318                            logistic_loss(pos_score, 1.0) + logistic_loss(neg_score, -1.0)
319                        }
320                        LossFunction::MarginRanking => {
321                            let margin = self
322                                .base
323                                .config
324                                .model_params
325                                .get("margin")
326                                .copied()
327                                .unwrap_or(1.0);
328                            margin_loss(pos_score, neg_score, margin)
329                        }
330                        LossFunction::SquaredLoss => (1.0 - pos_score).powi(2) + neg_score.powi(2),
331                    };
332
333                    batch_loss += triple_loss;
334
335                    // Compute and accumulate gradients
336                    let (entity_grads, relation_grads) =
337                        self.compute_gradients(pos_triple, neg_triple, pos_score, neg_score)?;
338
339                    batch_entity_grads += &entity_grads;
340                    batch_relation_grads += &relation_grads;
341                }
342            }
343
344            // Apply gradients with L2 regularization
345            gradient_update(
346                &mut self.entity_embeddings,
347                &batch_entity_grads,
348                learning_rate,
349                self.base.config.l2_reg,
350            );
351
352            gradient_update(
353                &mut self.relation_embeddings,
354                &batch_relation_grads,
355                learning_rate,
356                self.base.config.l2_reg,
357            );
358
359            // Optional: re-normalize entity embeddings
360            if self
361                .base
362                .config
363                .model_params
364                .get("normalize_entities")
365                .copied()
366                .unwrap_or(0.0)
367                > 0.0
368            {
369                normalize_embeddings(&mut self.entity_embeddings);
370            }
371
372            total_loss += batch_loss;
373        }
374
375        Ok(total_loss / num_batches as f64)
376    }
377}
378
379#[async_trait]
380impl EmbeddingModel for DistMult {
381    fn config(&self) -> &ModelConfig {
382        &self.base.config
383    }
384
385    fn model_id(&self) -> &Uuid {
386        &self.base.model_id
387    }
388
389    fn model_type(&self) -> &'static str {
390        "DistMult"
391    }
392
393    fn add_triple(&mut self, triple: Triple) -> Result<()> {
394        // Check for asymmetric relations and warn
395        let predicate_str = triple.predicate.to_string();
396        if predicate_str.contains("parent")
397            || predicate_str.contains("child")
398            || predicate_str.contains("born")
399            || predicate_str.contains("founder")
400        {
401            warn!(
402                "DistMult may not handle asymmetric relation well: {}",
403                predicate_str
404            );
405        }
406
407        self.base.add_triple(triple)
408    }
409
410    async fn train(&mut self, epochs: Option<usize>) -> Result<TrainingStats> {
411        let start_time = Instant::now();
412        let max_epochs = epochs.unwrap_or(self.base.config.max_epochs);
413
414        // Initialize embeddings if needed
415        self.initialize_embeddings();
416
417        if !self.embeddings_initialized {
418            return Err(anyhow!("No training data available"));
419        }
420
421        let mut loss_history = Vec::new();
422        let learning_rate = self.base.config.learning_rate;
423
424        info!("Starting DistMult training for {} epochs", max_epochs);
425
426        for epoch in 0..max_epochs {
427            let epoch_loss = self.train_epoch(learning_rate).await?;
428            loss_history.push(epoch_loss);
429
430            if epoch % 100 == 0 {
431                debug!("Epoch {}: loss = {:.6}", epoch, epoch_loss);
432            }
433
434            // Simple convergence check
435            if epoch > 10 && epoch_loss < 1e-6 {
436                info!("Converged at epoch {} with loss {:.6}", epoch, epoch_loss);
437                break;
438            }
439        }
440
441        self.base.mark_trained();
442        let training_time = start_time.elapsed().as_secs_f64();
443
444        Ok(TrainingStats {
445            epochs_completed: loss_history.len(),
446            final_loss: loss_history.last().copied().unwrap_or(0.0),
447            training_time_seconds: training_time,
448            convergence_achieved: loss_history.last().copied().unwrap_or(f64::INFINITY) < 1e-6,
449            loss_history,
450        })
451    }
452
453    fn get_entity_embedding(&self, entity: &str) -> Result<Vector> {
454        if !self.embeddings_initialized {
455            return Err(anyhow!("Model not trained"));
456        }
457
458        let entity_id = self
459            .base
460            .get_entity_id(entity)
461            .ok_or_else(|| anyhow!("Entity not found: {}", entity))?;
462
463        let embedding = self.entity_embeddings.row(entity_id).to_owned();
464        Ok(ndarray_to_vector(&embedding))
465    }
466
467    fn get_relation_embedding(&self, relation: &str) -> Result<Vector> {
468        if !self.embeddings_initialized {
469            return Err(anyhow!("Model not trained"));
470        }
471
472        let relation_id = self
473            .base
474            .get_relation_id(relation)
475            .ok_or_else(|| anyhow!("Relation not found: {}", relation))?;
476
477        let embedding = self.relation_embeddings.row(relation_id).to_owned();
478        Ok(ndarray_to_vector(&embedding))
479    }
480
481    fn score_triple(&self, subject: &str, predicate: &str, object: &str) -> Result<f64> {
482        let subject_id = self
483            .base
484            .get_entity_id(subject)
485            .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
486        let predicate_id = self
487            .base
488            .get_relation_id(predicate)
489            .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
490        let object_id = self
491            .base
492            .get_entity_id(object)
493            .ok_or_else(|| anyhow!("Object not found: {}", object))?;
494
495        self.score_triple_ids(subject_id, predicate_id, object_id)
496    }
497
498    fn predict_objects(
499        &self,
500        subject: &str,
501        predicate: &str,
502        k: usize,
503    ) -> Result<Vec<(String, f64)>> {
504        if !self.embeddings_initialized {
505            return Err(anyhow!("Model not trained"));
506        }
507
508        let subject_id = self
509            .base
510            .get_entity_id(subject)
511            .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
512        let predicate_id = self
513            .base
514            .get_relation_id(predicate)
515            .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
516
517        let mut scores = Vec::new();
518
519        for object_id in 0..self.base.num_entities() {
520            let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
521            let object_name = self
522                .base
523                .get_entity(object_id)
524                .expect("entity should exist for valid id")
525                .clone();
526            scores.push((object_name, score));
527        }
528
529        scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
530        scores.truncate(k);
531
532        Ok(scores)
533    }
534
535    fn predict_subjects(
536        &self,
537        predicate: &str,
538        object: &str,
539        k: usize,
540    ) -> Result<Vec<(String, f64)>> {
541        if !self.embeddings_initialized {
542            return Err(anyhow!("Model not trained"));
543        }
544
545        let predicate_id = self
546            .base
547            .get_relation_id(predicate)
548            .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
549        let object_id = self
550            .base
551            .get_entity_id(object)
552            .ok_or_else(|| anyhow!("Object not found: {}", object))?;
553
554        let mut scores = Vec::new();
555
556        for subject_id in 0..self.base.num_entities() {
557            let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
558            let subject_name = self
559                .base
560                .get_entity(subject_id)
561                .expect("entity should exist for valid id")
562                .clone();
563            scores.push((subject_name, score));
564        }
565
566        scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
567        scores.truncate(k);
568
569        Ok(scores)
570    }
571
572    fn predict_relations(
573        &self,
574        subject: &str,
575        object: &str,
576        k: usize,
577    ) -> Result<Vec<(String, f64)>> {
578        if !self.embeddings_initialized {
579            return Err(anyhow!("Model not trained"));
580        }
581
582        let subject_id = self
583            .base
584            .get_entity_id(subject)
585            .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
586        let object_id = self
587            .base
588            .get_entity_id(object)
589            .ok_or_else(|| anyhow!("Object not found: {}", object))?;
590
591        let mut scores = Vec::new();
592
593        for predicate_id in 0..self.base.num_relations() {
594            let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
595            let predicate_name = self
596                .base
597                .get_relation(predicate_id)
598                .expect("relation should exist for valid id")
599                .clone();
600            scores.push((predicate_name, score));
601        }
602
603        scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
604        scores.truncate(k);
605
606        Ok(scores)
607    }
608
609    fn get_entities(&self) -> Vec<String> {
610        self.base.get_entities()
611    }
612
613    fn get_relations(&self) -> Vec<String> {
614        self.base.get_relations()
615    }
616
617    fn get_stats(&self) -> ModelStats {
618        self.base.get_stats("DistMult")
619    }
620
621    fn save(&self, path: &str) -> Result<()> {
622        info!("Saving DistMult model to {}", path);
623
624        let serializable = DistMultSerializable {
625            base: BaseModelSnapshot::capture(&self.base),
626            entity_embeddings: MatrixF64::from_array(&self.entity_embeddings),
627            relation_embeddings: MatrixF64::from_array(&self.relation_embeddings),
628            embeddings_initialized: self.embeddings_initialized,
629            dropout_rate: self.dropout_rate,
630            loss_function: self.loss_function,
631        };
632
633        let file = File::create(path)
634            .map_err(|e| anyhow!("Failed to create model file {}: {}", path, e))?;
635        let writer = BufWriter::new(file);
636        oxicode::serde::encode_into_std_write(&serializable, writer, oxicode::config::standard())
637            .map_err(|e| anyhow!("Failed to serialize DistMult model: {}", e))?;
638
639        info!("DistMult model saved successfully");
640        Ok(())
641    }
642
643    fn load(&mut self, path: &str) -> Result<()> {
644        info!("Loading DistMult model from {}", path);
645
646        if !Path::new(path).exists() {
647            return Err(anyhow!("Model file not found: {}", path));
648        }
649
650        let file =
651            File::open(path).map_err(|e| anyhow!("Failed to open model file {}: {}", path, e))?;
652        let reader = BufReader::new(file);
653        let (serializable, _): (DistMultSerializable, _) =
654            oxicode::serde::decode_from_std_read(reader, oxicode::config::standard())
655                .map_err(|e| anyhow!("Failed to deserialize DistMult model: {}", e))?;
656
657        self.entity_embeddings = serializable.entity_embeddings.to_array()?;
658        self.relation_embeddings = serializable.relation_embeddings.to_array()?;
659        self.embeddings_initialized = serializable.embeddings_initialized;
660        self.dropout_rate = serializable.dropout_rate;
661        self.loss_function = serializable.loss_function;
662        serializable.base.restore_into(&mut self.base);
663
664        info!("DistMult model loaded successfully");
665        Ok(())
666    }
667
668    fn clear(&mut self) {
669        self.base.clear();
670        self.entity_embeddings = Array2::zeros((0, self.base.config.dimensions));
671        self.relation_embeddings = Array2::zeros((0, self.base.config.dimensions));
672        self.embeddings_initialized = false;
673    }
674
675    fn is_trained(&self) -> bool {
676        self.base.is_trained
677    }
678
679    async fn encode(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>> {
680        Err(anyhow!(
681            "Knowledge graph embedding model does not support text encoding"
682        ))
683    }
684}
685
686#[cfg(test)]
687mod tests {
688    use super::*;
689    use crate::NamedNode;
690
691    #[tokio::test]
692    async fn test_distmult_basic() -> Result<()> {
693        let config = ModelConfig::default()
694            .with_dimensions(50)
695            .with_max_epochs(10)
696            .with_seed(42);
697
698        let mut model = DistMult::new(config);
699
700        // Add test triples (use symmetric relation for DistMult)
701        let alice = NamedNode::new("http://example.org/alice")?;
702        let similar_to = NamedNode::new("http://example.org/similarTo")?;
703        let bob = NamedNode::new("http://example.org/bob")?;
704
705        model.add_triple(Triple::new(alice.clone(), similar_to.clone(), bob.clone()))?;
706        model.add_triple(Triple::new(bob.clone(), similar_to.clone(), alice.clone()))?;
707
708        // Train
709        let stats = model.train(Some(5)).await?;
710        assert!(stats.epochs_completed > 0);
711
712        // Test embeddings
713        let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
714        assert_eq!(alice_emb.dimensions, 50);
715
716        // Test scoring
717        let score = model.score_triple(
718            "http://example.org/alice",
719            "http://example.org/similarTo",
720            "http://example.org/bob",
721        )?;
722
723        // Score should be a finite number
724        assert!(score.is_finite());
725
726        Ok(())
727    }
728}