Skip to main content

oxirs_embed/models/
tucker.rs

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