Skip to main content

oxirs_embed/models/
complex.rs

1//! ComplEx: Complex Embeddings for Simple Link Prediction
2//!
3//! ComplEx uses complex-valued embeddings to better model asymmetric relations.
4//! The scoring function is: Re(<h, r, conj(t)>) where Re denotes real part,
5//! <> denotes complex dot product, and conj denotes complex conjugate.
6//!
7//! Reference: Trouillon et al. "Complex Embeddings for Simple Link Prediction" (2016)
8
9use 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/// Serializable representation of a ComplEx model for persistence.
27#[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
38/// Type alias for gradient tensors
39type GradientTuple = (Array2<f64>, Array2<f64>, Array2<f64>, Array2<f64>);
40
41/// ComplEx embedding model using complex-valued embeddings
42#[derive(Debug)]
43pub struct ComplEx {
44    /// Base model functionality
45    base: BaseModel,
46    /// Real part of entity embeddings (num_entities × dimensions)
47    entity_embeddings_real: Array2<f64>,
48    /// Imaginary part of entity embeddings (num_entities × dimensions)
49    entity_embeddings_imag: Array2<f64>,
50    /// Real part of relation embeddings (num_relations × dimensions)
51    relation_embeddings_real: Array2<f64>,
52    /// Imaginary part of relation embeddings (num_relations × dimensions)
53    relation_embeddings_imag: Array2<f64>,
54    /// Whether embeddings have been initialized
55    embeddings_initialized: bool,
56    /// Regularization method
57    regularization: RegularizationType,
58}
59
60/// Regularization types for ComplEx
61#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
62pub enum RegularizationType {
63    /// L2 regularization on embeddings
64    L2,
65    /// N3 regularization (nuclear 3-norm)
66    N3,
67    /// No additional regularization
68    None,
69}
70
71impl ComplEx {
72    /// Create a new ComplEx model
73    pub fn new(config: ModelConfig) -> Self {
74        let base = BaseModel::new(config.clone());
75
76        // Get ComplEx-specific parameters
77        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, // Default to N3
82        };
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    /// Initialize complex embeddings
96    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        // Initialize all embedding components with Xavier initialization
112        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    /// Score a triple using ComplEx scoring function
140    /// Score = Re(<h, r, conj(t)>) = Re(h) * Re(r) * Re(t) + Re(h) * Im(r) * Im(t) +
141    ///                                 Im(h) * Re(r) * Im(t) - Im(h) * Im(r) * Re(t)
142    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        // Complex multiplication: (h_real + i*h_imag) * (r_real + i*r_imag) * conj(t_real + i*t_imag)
160        // = (h_real + i*h_imag) * (r_real + i*r_imag) * (t_real - i*t_imag)
161        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    /// Compute gradients for ComplEx model
170    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        // Logistic loss gradients
183        let pos_sigmoid = sigmoid(pos_score);
184        let neg_sigmoid = sigmoid(neg_score);
185
186        let pos_grad_coeff = pos_sigmoid - 1.0; // Derivative of log(sigmoid(x))
187        let neg_grad_coeff = neg_sigmoid; // Derivative of log(1 - sigmoid(x))
188
189        // Compute gradients for positive triple
190        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        // Compute gradients for negative triple
200        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    /// Add gradients for a single triple
218    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        // Gradients w.r.t. h (subject)
237        // ∂score/∂h_real = r_real * t_real + r_imag * t_imag
238        // ∂score/∂h_imag = r_real * t_imag - r_imag * t_real
239        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        // Gradients w.r.t. r (relation)
246        // ∂score/∂r_real = h_real * t_real + h_imag * t_imag
247        // ∂score/∂r_imag = h_real * t_imag - h_imag * t_real
248        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        // Gradients w.r.t. t (object) - note the conjugate
255        // ∂score/∂t_real = h_real * r_real - h_imag * r_imag
256        // ∂score/∂t_imag = -(h_real * r_imag + h_imag * r_real)
257        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    /// Apply N3 regularization
265    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        // N3 regularization: penalize the nuclear 3-norm
274        // For complex embeddings, this becomes more involved
275        // For simplicity, we apply L2 regularization here
276        // A full N3 implementation would require more complex tensor operations
277
278        *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    /// Perform one training epoch
285    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        // Create shuffled batches
293        let mut shuffled_triples = self.base.triples.clone();
294        // Manual Fisher-Yates shuffle using scirs2-core
295        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                // Generate negative samples
311                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                    // Compute scores
317                    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                    // Compute logistic loss
323                    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                    // Compute and accumulate gradients
330                    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            // Apply regularization
345            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            // Apply gradients
366            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    /// Get entity embedding as a concatenated real/imaginary vector
378    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        // Concatenate real and imaginary parts
383        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    /// Get relation embedding as a concatenated real/imaginary vector
395    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        // Concatenate real and imaginary parts
400        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        // Initialize embeddings if needed
435        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            // Simple convergence check
455            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        // Add test triples
723        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        // Train
730        let stats = model.train(Some(5)).await?;
731        assert!(stats.epochs_completed > 0);
732
733        // Test embeddings (should be 2x dimensions due to complex)
734        let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
735        assert_eq!(alice_emb.dimensions, 100); // 2 * 50
736
737        // Test scoring
738        let score = model.score_triple(
739            "http://example.org/alice",
740            "http://example.org/knows",
741            "http://example.org/bob",
742        )?;
743
744        // Score should be a finite number
745        assert!(score.is_finite());
746
747        Ok(())
748    }
749}