1use crate::models::serialization::{BaseModelSnapshot, MatrixF64};
10use crate::models::{common::*, BaseModel};
11use crate::{EmbeddingModel, ModelConfig, ModelStats, TrainingStats, Triple, Vector};
12use anyhow::{anyhow, Result};
13use async_trait::async_trait;
14use scirs2_core::ndarray_ext::Array2;
15#[allow(unused_imports)]
16use scirs2_core::random::{Random, RngExt};
17use serde::{Deserialize, Serialize};
18use std::fs::File;
19use std::io::{BufReader, BufWriter};
20use std::ops::AddAssign;
21use std::path::Path;
22use std::time::Instant;
23use tracing::{debug, info};
24use uuid::Uuid;
25
26#[derive(Debug, Serialize, Deserialize)]
28struct ComplExSerializable {
29 base: BaseModelSnapshot,
30 entity_embeddings_real: MatrixF64,
31 entity_embeddings_imag: MatrixF64,
32 relation_embeddings_real: MatrixF64,
33 relation_embeddings_imag: MatrixF64,
34 embeddings_initialized: bool,
35 regularization: RegularizationType,
36}
37
38type GradientTuple = (Array2<f64>, Array2<f64>, Array2<f64>, Array2<f64>);
40
41#[derive(Debug)]
43pub struct ComplEx {
44 base: BaseModel,
46 entity_embeddings_real: Array2<f64>,
48 entity_embeddings_imag: Array2<f64>,
50 relation_embeddings_real: Array2<f64>,
52 relation_embeddings_imag: Array2<f64>,
54 embeddings_initialized: bool,
56 regularization: RegularizationType,
58}
59
60#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
62pub enum RegularizationType {
63 L2,
65 N3,
67 None,
69}
70
71impl ComplEx {
72 pub fn new(config: ModelConfig) -> Self {
74 let base = BaseModel::new(config.clone());
75
76 let regularization = match config.model_params.get("regularization") {
78 Some(0.0) => RegularizationType::None,
79 Some(1.0) => RegularizationType::L2,
80 Some(2.0) => RegularizationType::N3,
81 _ => RegularizationType::N3, };
83
84 Self {
85 base,
86 entity_embeddings_real: Array2::zeros((0, config.dimensions)),
87 entity_embeddings_imag: Array2::zeros((0, config.dimensions)),
88 relation_embeddings_real: Array2::zeros((0, config.dimensions)),
89 relation_embeddings_imag: Array2::zeros((0, config.dimensions)),
90 embeddings_initialized: false,
91 regularization,
92 }
93 }
94
95 fn initialize_embeddings(&mut self) {
97 if self.embeddings_initialized {
98 return;
99 }
100
101 let num_entities = self.base.num_entities();
102 let num_relations = self.base.num_relations();
103 let dimensions = self.base.config.dimensions;
104
105 if num_entities == 0 || num_relations == 0 {
106 return;
107 }
108
109 let mut rng = Random::default();
110
111 self.entity_embeddings_real =
113 xavier_init((num_entities, dimensions), dimensions, dimensions, &mut rng);
114
115 self.entity_embeddings_imag =
116 xavier_init((num_entities, dimensions), dimensions, dimensions, &mut rng);
117
118 self.relation_embeddings_real = xavier_init(
119 (num_relations, dimensions),
120 dimensions,
121 dimensions,
122 &mut rng,
123 );
124
125 self.relation_embeddings_imag = xavier_init(
126 (num_relations, dimensions),
127 dimensions,
128 dimensions,
129 &mut rng,
130 );
131
132 self.embeddings_initialized = true;
133 debug!(
134 "Initialized ComplEx embeddings: {} entities, {} relations, {} dimensions",
135 num_entities, num_relations, dimensions
136 );
137 }
138
139 fn score_triple_ids(
143 &self,
144 subject_id: usize,
145 predicate_id: usize,
146 object_id: usize,
147 ) -> Result<f64> {
148 if !self.embeddings_initialized {
149 return Err(anyhow!("Model not trained"));
150 }
151
152 let h_real = self.entity_embeddings_real.row(subject_id);
153 let h_imag = self.entity_embeddings_imag.row(subject_id);
154 let r_real = self.relation_embeddings_real.row(predicate_id);
155 let r_imag = self.relation_embeddings_imag.row(predicate_id);
156 let t_real = self.entity_embeddings_real.row(object_id);
157 let t_imag = self.entity_embeddings_imag.row(object_id);
158
159 let score = (&h_real * &r_real * t_real).sum()
162 + (&h_real * &r_imag * t_imag).sum()
163 + (&h_imag * &r_real * t_imag).sum()
164 - (&h_imag * &r_imag * t_real).sum();
165
166 Ok(score)
167 }
168
169 fn compute_gradients(
171 &self,
172 pos_triple: (usize, usize, usize),
173 neg_triple: (usize, usize, usize),
174 pos_score: f64,
175 neg_score: f64,
176 ) -> Result<GradientTuple> {
177 let mut entity_grads_real = Array2::zeros(self.entity_embeddings_real.raw_dim());
178 let mut entity_grads_imag = Array2::zeros(self.entity_embeddings_imag.raw_dim());
179 let mut relation_grads_real = Array2::zeros(self.relation_embeddings_real.raw_dim());
180 let mut relation_grads_imag = Array2::zeros(self.relation_embeddings_imag.raw_dim());
181
182 let pos_sigmoid = sigmoid(pos_score);
184 let neg_sigmoid = sigmoid(neg_score);
185
186 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_real,
194 &mut entity_grads_imag,
195 &mut relation_grads_real,
196 &mut relation_grads_imag,
197 );
198
199 self.add_triple_gradients(
201 neg_triple,
202 neg_grad_coeff,
203 &mut entity_grads_real,
204 &mut entity_grads_imag,
205 &mut relation_grads_real,
206 &mut relation_grads_imag,
207 );
208
209 Ok((
210 entity_grads_real,
211 entity_grads_imag,
212 relation_grads_real,
213 relation_grads_imag,
214 ))
215 }
216
217 fn add_triple_gradients(
219 &self,
220 triple: (usize, usize, usize),
221 grad_coeff: f64,
222 entity_grads_real: &mut Array2<f64>,
223 entity_grads_imag: &mut Array2<f64>,
224 relation_grads_real: &mut Array2<f64>,
225 relation_grads_imag: &mut Array2<f64>,
226 ) {
227 let (s, p, o) = triple;
228
229 let h_real = self.entity_embeddings_real.row(s);
230 let h_imag = self.entity_embeddings_imag.row(s);
231 let r_real = self.relation_embeddings_real.row(p);
232 let r_imag = self.relation_embeddings_imag.row(p);
233 let t_real = self.entity_embeddings_real.row(o);
234 let t_imag = self.entity_embeddings_imag.row(o);
235
236 let h_real_grad = (&r_real * &t_real + &r_imag * &t_imag) * grad_coeff;
240 let h_imag_grad = (&r_real * &t_imag - &r_imag * &t_real) * grad_coeff;
241
242 entity_grads_real.row_mut(s).add_assign(&h_real_grad);
243 entity_grads_imag.row_mut(s).add_assign(&h_imag_grad);
244
245 let r_real_grad = (&h_real * &t_real + &h_imag * &t_imag) * grad_coeff;
249 let r_imag_grad = (&h_real * &t_imag - &h_imag * &t_real) * grad_coeff;
250
251 relation_grads_real.row_mut(p).add_assign(&r_real_grad);
252 relation_grads_imag.row_mut(p).add_assign(&r_imag_grad);
253
254 let t_real_grad = (&h_real * &r_real - &h_imag * &r_imag) * grad_coeff;
258 let t_imag_grad = -(&h_real * &r_imag + &h_imag * &r_real) * grad_coeff;
259
260 entity_grads_real.row_mut(o).add_assign(&t_real_grad);
261 entity_grads_imag.row_mut(o).add_assign(&t_imag_grad);
262 }
263
264 fn apply_n3_regularization(
266 &self,
267 entity_grads_real: &mut Array2<f64>,
268 entity_grads_imag: &mut Array2<f64>,
269 relation_grads_real: &mut Array2<f64>,
270 relation_grads_imag: &mut Array2<f64>,
271 regularization_weight: f64,
272 ) {
273 *entity_grads_real += &(&self.entity_embeddings_real * regularization_weight);
279 *entity_grads_imag += &(&self.entity_embeddings_imag * regularization_weight);
280 *relation_grads_real += &(&self.relation_embeddings_real * regularization_weight);
281 *relation_grads_imag += &(&self.relation_embeddings_imag * regularization_weight);
282 }
283
284 async fn train_epoch(&mut self, learning_rate: f64) -> Result<f64> {
286 let mut rng = Random::default();
287
288 let mut total_loss = 0.0;
289 let num_batches = (self.base.triples.len() + self.base.config.batch_size - 1)
290 / self.base.config.batch_size;
291
292 let mut shuffled_triples = self.base.triples.clone();
294 for i in (1..shuffled_triples.len()).rev() {
296 let j = rng.random_range(0..i + 1);
297 shuffled_triples.swap(i, j);
298 }
299
300 for batch_triples in shuffled_triples.chunks(self.base.config.batch_size) {
301 let mut batch_entity_grads_real = Array2::zeros(self.entity_embeddings_real.raw_dim());
302 let mut batch_entity_grads_imag = Array2::zeros(self.entity_embeddings_imag.raw_dim());
303 let mut batch_relation_grads_real =
304 Array2::zeros(self.relation_embeddings_real.raw_dim());
305 let mut batch_relation_grads_imag =
306 Array2::zeros(self.relation_embeddings_imag.raw_dim());
307 let mut batch_loss = 0.0;
308
309 for &pos_triple in batch_triples {
310 let neg_samples = self
312 .base
313 .generate_negative_samples(self.base.config.negative_samples, &mut rng);
314
315 for neg_triple in neg_samples {
316 let pos_score =
318 self.score_triple_ids(pos_triple.0, pos_triple.1, pos_triple.2)?;
319 let neg_score =
320 self.score_triple_ids(neg_triple.0, neg_triple.1, neg_triple.2)?;
321
322 let pos_loss = logistic_loss(pos_score, 1.0);
324 let neg_loss = logistic_loss(neg_score, -1.0);
325 let total_triple_loss = pos_loss + neg_loss;
326
327 batch_loss += total_triple_loss;
328
329 let (
331 entity_grads_real,
332 entity_grads_imag,
333 relation_grads_real,
334 relation_grads_imag,
335 ) = self.compute_gradients(pos_triple, neg_triple, pos_score, neg_score)?;
336
337 batch_entity_grads_real += &entity_grads_real;
338 batch_entity_grads_imag += &entity_grads_imag;
339 batch_relation_grads_real += &relation_grads_real;
340 batch_relation_grads_imag += &relation_grads_imag;
341 }
342 }
343
344 match self.regularization {
346 RegularizationType::L2 => {
347 let reg_weight = self.base.config.l2_reg;
348 batch_entity_grads_real += &(&self.entity_embeddings_real * reg_weight);
349 batch_entity_grads_imag += &(&self.entity_embeddings_imag * reg_weight);
350 batch_relation_grads_real += &(&self.relation_embeddings_real * reg_weight);
351 batch_relation_grads_imag += &(&self.relation_embeddings_imag * reg_weight);
352 }
353 RegularizationType::N3 => {
354 self.apply_n3_regularization(
355 &mut batch_entity_grads_real,
356 &mut batch_entity_grads_imag,
357 &mut batch_relation_grads_real,
358 &mut batch_relation_grads_imag,
359 self.base.config.l2_reg,
360 );
361 }
362 RegularizationType::None => {}
363 }
364
365 self.entity_embeddings_real -= &(&batch_entity_grads_real * learning_rate);
367 self.entity_embeddings_imag -= &(&batch_entity_grads_imag * learning_rate);
368 self.relation_embeddings_real -= &(&batch_relation_grads_real * learning_rate);
369 self.relation_embeddings_imag -= &(&batch_relation_grads_imag * learning_rate);
370
371 total_loss += batch_loss;
372 }
373
374 Ok(total_loss / num_batches as f64)
375 }
376
377 fn get_entity_embedding_vector(&self, entity_id: usize) -> Vector {
379 let real_part = self.entity_embeddings_real.row(entity_id);
380 let imag_part = self.entity_embeddings_imag.row(entity_id);
381
382 let mut values = Vec::with_capacity(real_part.len() * 2);
384 for &val in real_part.iter() {
385 values.push(val as f32);
386 }
387 for &val in imag_part.iter() {
388 values.push(val as f32);
389 }
390
391 Vector::new(values)
392 }
393
394 fn get_relation_embedding_vector(&self, relation_id: usize) -> Vector {
396 let real_part = self.relation_embeddings_real.row(relation_id);
397 let imag_part = self.relation_embeddings_imag.row(relation_id);
398
399 let mut values = Vec::with_capacity(real_part.len() * 2);
401 for &val in real_part.iter() {
402 values.push(val as f32);
403 }
404 for &val in imag_part.iter() {
405 values.push(val as f32);
406 }
407
408 Vector::new(values)
409 }
410}
411
412#[async_trait]
413impl EmbeddingModel for ComplEx {
414 fn config(&self) -> &ModelConfig {
415 &self.base.config
416 }
417
418 fn model_id(&self) -> &Uuid {
419 &self.base.model_id
420 }
421
422 fn model_type(&self) -> &'static str {
423 "ComplEx"
424 }
425
426 fn add_triple(&mut self, triple: Triple) -> Result<()> {
427 self.base.add_triple(triple)
428 }
429
430 async fn train(&mut self, epochs: Option<usize>) -> Result<TrainingStats> {
431 let start_time = Instant::now();
432 let max_epochs = epochs.unwrap_or(self.base.config.max_epochs);
433
434 self.initialize_embeddings();
436
437 if !self.embeddings_initialized {
438 return Err(anyhow!("No training data available"));
439 }
440
441 let mut loss_history = Vec::new();
442 let learning_rate = self.base.config.learning_rate;
443
444 info!("Starting ComplEx training for {} epochs", max_epochs);
445
446 for epoch in 0..max_epochs {
447 let epoch_loss = self.train_epoch(learning_rate).await?;
448 loss_history.push(epoch_loss);
449
450 if epoch % 100 == 0 {
451 debug!("Epoch {}: loss = {:.6}", epoch, epoch_loss);
452 }
453
454 if epoch > 10 && epoch_loss < 1e-6 {
456 info!("Converged at epoch {} with loss {:.6}", epoch, epoch_loss);
457 break;
458 }
459 }
460
461 self.base.mark_trained();
462 let training_time = start_time.elapsed().as_secs_f64();
463
464 Ok(TrainingStats {
465 epochs_completed: loss_history.len(),
466 final_loss: loss_history.last().copied().unwrap_or(0.0),
467 training_time_seconds: training_time,
468 convergence_achieved: loss_history.last().copied().unwrap_or(f64::INFINITY) < 1e-6,
469 loss_history,
470 })
471 }
472
473 fn get_entity_embedding(&self, entity: &str) -> Result<Vector> {
474 if !self.embeddings_initialized {
475 return Err(anyhow!("Model not trained"));
476 }
477
478 let entity_id = self
479 .base
480 .get_entity_id(entity)
481 .ok_or_else(|| anyhow!("Entity not found: {}", entity))?;
482
483 Ok(self.get_entity_embedding_vector(entity_id))
484 }
485
486 fn get_relation_embedding(&self, relation: &str) -> Result<Vector> {
487 if !self.embeddings_initialized {
488 return Err(anyhow!("Model not trained"));
489 }
490
491 let relation_id = self
492 .base
493 .get_relation_id(relation)
494 .ok_or_else(|| anyhow!("Relation not found: {}", relation))?;
495
496 Ok(self.get_relation_embedding_vector(relation_id))
497 }
498
499 fn score_triple(&self, subject: &str, predicate: &str, object: &str) -> Result<f64> {
500 let subject_id = self
501 .base
502 .get_entity_id(subject)
503 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
504 let predicate_id = self
505 .base
506 .get_relation_id(predicate)
507 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
508 let object_id = self
509 .base
510 .get_entity_id(object)
511 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
512
513 self.score_triple_ids(subject_id, predicate_id, object_id)
514 }
515
516 fn predict_objects(
517 &self,
518 subject: &str,
519 predicate: &str,
520 k: usize,
521 ) -> Result<Vec<(String, f64)>> {
522 if !self.embeddings_initialized {
523 return Err(anyhow!("Model not trained"));
524 }
525
526 let subject_id = self
527 .base
528 .get_entity_id(subject)
529 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
530 let predicate_id = self
531 .base
532 .get_relation_id(predicate)
533 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
534
535 let mut scores = Vec::new();
536
537 for object_id in 0..self.base.num_entities() {
538 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
539 let object_name = self
540 .base
541 .get_entity(object_id)
542 .expect("entity should exist for valid id")
543 .clone();
544 scores.push((object_name, score));
545 }
546
547 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
548 scores.truncate(k);
549
550 Ok(scores)
551 }
552
553 fn predict_subjects(
554 &self,
555 predicate: &str,
556 object: &str,
557 k: usize,
558 ) -> Result<Vec<(String, f64)>> {
559 if !self.embeddings_initialized {
560 return Err(anyhow!("Model not trained"));
561 }
562
563 let predicate_id = self
564 .base
565 .get_relation_id(predicate)
566 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
567 let object_id = self
568 .base
569 .get_entity_id(object)
570 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
571
572 let mut scores = Vec::new();
573
574 for subject_id in 0..self.base.num_entities() {
575 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
576 let subject_name = self
577 .base
578 .get_entity(subject_id)
579 .expect("entity should exist for valid id")
580 .clone();
581 scores.push((subject_name, score));
582 }
583
584 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
585 scores.truncate(k);
586
587 Ok(scores)
588 }
589
590 fn predict_relations(
591 &self,
592 subject: &str,
593 object: &str,
594 k: usize,
595 ) -> Result<Vec<(String, f64)>> {
596 if !self.embeddings_initialized {
597 return Err(anyhow!("Model not trained"));
598 }
599
600 let subject_id = self
601 .base
602 .get_entity_id(subject)
603 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
604 let object_id = self
605 .base
606 .get_entity_id(object)
607 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
608
609 let mut scores = Vec::new();
610
611 for predicate_id in 0..self.base.num_relations() {
612 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
613 let predicate_name = self
614 .base
615 .get_relation(predicate_id)
616 .expect("relation should exist for valid id")
617 .clone();
618 scores.push((predicate_name, score));
619 }
620
621 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
622 scores.truncate(k);
623
624 Ok(scores)
625 }
626
627 fn get_entities(&self) -> Vec<String> {
628 self.base.get_entities()
629 }
630
631 fn get_relations(&self) -> Vec<String> {
632 self.base.get_relations()
633 }
634
635 fn get_stats(&self) -> ModelStats {
636 self.base.get_stats("ComplEx")
637 }
638
639 fn save(&self, path: &str) -> Result<()> {
640 info!("Saving ComplEx model to {}", path);
641
642 let serializable = ComplExSerializable {
643 base: BaseModelSnapshot::capture(&self.base),
644 entity_embeddings_real: MatrixF64::from_array(&self.entity_embeddings_real),
645 entity_embeddings_imag: MatrixF64::from_array(&self.entity_embeddings_imag),
646 relation_embeddings_real: MatrixF64::from_array(&self.relation_embeddings_real),
647 relation_embeddings_imag: MatrixF64::from_array(&self.relation_embeddings_imag),
648 embeddings_initialized: self.embeddings_initialized,
649 regularization: self.regularization,
650 };
651
652 let file = File::create(path)
653 .map_err(|e| anyhow!("Failed to create model file {}: {}", path, e))?;
654 let writer = BufWriter::new(file);
655 oxicode::serde::encode_into_std_write(&serializable, writer, oxicode::config::standard())
656 .map_err(|e| anyhow!("Failed to serialize ComplEx model: {}", e))?;
657
658 info!("ComplEx model saved successfully");
659 Ok(())
660 }
661
662 fn load(&mut self, path: &str) -> Result<()> {
663 info!("Loading ComplEx model from {}", path);
664
665 if !Path::new(path).exists() {
666 return Err(anyhow!("Model file not found: {}", path));
667 }
668
669 let file =
670 File::open(path).map_err(|e| anyhow!("Failed to open model file {}: {}", path, e))?;
671 let reader = BufReader::new(file);
672 let (serializable, _): (ComplExSerializable, _) =
673 oxicode::serde::decode_from_std_read(reader, oxicode::config::standard())
674 .map_err(|e| anyhow!("Failed to deserialize ComplEx model: {}", e))?;
675
676 self.entity_embeddings_real = serializable.entity_embeddings_real.to_array()?;
677 self.entity_embeddings_imag = serializable.entity_embeddings_imag.to_array()?;
678 self.relation_embeddings_real = serializable.relation_embeddings_real.to_array()?;
679 self.relation_embeddings_imag = serializable.relation_embeddings_imag.to_array()?;
680 self.embeddings_initialized = serializable.embeddings_initialized;
681 self.regularization = serializable.regularization;
682 serializable.base.restore_into(&mut self.base);
683
684 info!("ComplEx model loaded successfully");
685 Ok(())
686 }
687
688 fn clear(&mut self) {
689 self.base.clear();
690 self.entity_embeddings_real = Array2::zeros((0, self.base.config.dimensions));
691 self.entity_embeddings_imag = Array2::zeros((0, self.base.config.dimensions));
692 self.relation_embeddings_real = Array2::zeros((0, self.base.config.dimensions));
693 self.relation_embeddings_imag = Array2::zeros((0, self.base.config.dimensions));
694 self.embeddings_initialized = false;
695 }
696
697 fn is_trained(&self) -> bool {
698 self.base.is_trained
699 }
700
701 async fn encode(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>> {
702 Err(anyhow!(
703 "Knowledge graph embedding model does not support text encoding"
704 ))
705 }
706}
707
708#[cfg(test)]
709mod tests {
710 use super::*;
711 use crate::NamedNode;
712
713 #[tokio::test]
714 async fn test_complex_basic() -> Result<()> {
715 let config = ModelConfig::default()
716 .with_dimensions(50)
717 .with_max_epochs(10)
718 .with_seed(42);
719
720 let mut model = ComplEx::new(config);
721
722 let alice = NamedNode::new("http://example.org/alice")?;
724 let knows = NamedNode::new("http://example.org/knows")?;
725 let bob = NamedNode::new("http://example.org/bob")?;
726
727 model.add_triple(Triple::new(alice.clone(), knows.clone(), bob.clone()))?;
728
729 let stats = model.train(Some(5)).await?;
731 assert!(stats.epochs_completed > 0);
732
733 let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
735 assert_eq!(alice_emb.dimensions, 100); let score = model.score_triple(
739 "http://example.org/alice",
740 "http://example.org/knows",
741 "http://example.org/bob",
742 )?;
743
744 assert!(score.is_finite());
746
747 Ok(())
748 }
749}