1use 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#[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#[derive(Debug)]
40pub struct DistMult {
41 base: BaseModel,
43 entity_embeddings: Array2<f64>,
45 relation_embeddings: Array2<f64>,
47 embeddings_initialized: bool,
49 #[allow(dead_code)]
51 dropout_rate: f64,
52 loss_function: LossFunction,
54}
55
56#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
58pub enum LossFunction {
59 Logistic,
61 MarginRanking,
63 SquaredLoss,
65}
66
67impl DistMult {
68 pub fn new(config: ModelConfig) -> Self {
70 let base = BaseModel::new(config.clone());
71
72 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, };
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 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 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 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 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 let score = (&h * &r * t).sum();
150
151 Ok(score)
152 }
153
154 #[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 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 let pos_sigmoid = sigmoid(pos_score);
185 let neg_sigmoid = sigmoid(neg_score);
186
187 let pos_grad_coeff = pos_sigmoid - 1.0; let neg_grad_coeff = neg_sigmoid; 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 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 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 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 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 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 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 let mut shuffled_triples = self.base.triples.clone();
291 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 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 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 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 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 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 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 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 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 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 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 let stats = model.train(Some(5)).await?;
710 assert!(stats.epochs_completed > 0);
711
712 let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
714 assert_eq!(alice_emb.dimensions, 50);
715
716 let score = model.score_triple(
718 "http://example.org/alice",
719 "http://example.org/similarTo",
720 "http://example.org/bob",
721 )?;
722
723 assert!(score.is_finite());
725
726 Ok(())
727 }
728}