Skip to main content

oxirs_vec/
gnn_embeddings.rs

1//! Graph Neural Network (GNN) embeddings for knowledge graphs
2//!
3//! This module implements GNN-based embedding methods:
4//! - GCN: Graph Convolutional Networks
5//! - GraphSAGE: Graph Sample and Aggregate
6
7use crate::random_utils::NormalSampler as Normal;
8use crate::{
9    kg_embeddings::{KGEmbeddingConfig, KGEmbeddingModel, Triple},
10    Vector,
11};
12use anyhow::{anyhow, Result};
13use nalgebra::{DMatrix, DVector};
14use scirs2_core::random::{Random, Rng, RngExt};
15use std::collections::HashMap;
16
17/// Graph Convolutional Network (GCN) embedding model
18pub struct GCN {
19    config: KGEmbeddingConfig,
20    entity_embeddings: HashMap<String, DVector<f32>>,
21    relation_embeddings: HashMap<String, DVector<f32>>,
22    entities: Vec<String>,
23    relations: Vec<String>,
24    adjacency_matrix: Option<DMatrix<f32>>,
25    weight_matrices: Vec<DMatrix<f32>>,
26    num_layers: usize,
27}
28
29impl GCN {
30    pub fn new(config: KGEmbeddingConfig) -> Self {
31        let num_layers = 2; // Default to 2 layers
32        Self {
33            config,
34            entity_embeddings: HashMap::new(),
35            relation_embeddings: HashMap::new(),
36            entities: Vec::new(),
37            relations: Vec::new(),
38            adjacency_matrix: None,
39            weight_matrices: Vec::new(),
40            num_layers,
41        }
42    }
43
44    /// Initialize GCN with specified number of layers
45    pub fn with_layers(config: KGEmbeddingConfig, num_layers: usize) -> Self {
46        Self {
47            config,
48            entity_embeddings: HashMap::new(),
49            relation_embeddings: HashMap::new(),
50            entities: Vec::new(),
51            relations: Vec::new(),
52            adjacency_matrix: None,
53            weight_matrices: Vec::new(),
54            num_layers,
55        }
56    }
57
58    /// Initialize embeddings and graph structure
59    fn initialize(&mut self, triples: &[Triple]) -> Result<()> {
60        // Collect unique entities and relations
61        let mut entities = std::collections::HashSet::new();
62        let mut relations = std::collections::HashSet::new();
63
64        for triple in triples {
65            entities.insert(triple.subject.clone());
66            entities.insert(triple.object.clone());
67            relations.insert(triple.predicate.clone());
68        }
69
70        self.entities = entities.into_iter().collect();
71        self.relations = relations.into_iter().collect();
72
73        let _num_entities = self.entities.len();
74
75        // Initialize entity embeddings
76        let mut rng = if let Some(seed) = self.config.random_seed {
77            Random::seed(seed)
78        } else {
79            Random::seed(42)
80        };
81
82        let normal = Normal::new(0.0, 0.1)
83            .map_err(|e| anyhow!("Failed to create normal distribution: {}", e))?;
84
85        for entity in &self.entities {
86            let embedding: Vec<f32> = (0..self.config.dimensions)
87                .map(|_| normal.sample(&mut rng))
88                .collect();
89            self.entity_embeddings
90                .insert(entity.clone(), DVector::from_vec(embedding));
91        }
92
93        for relation in &self.relations {
94            let embedding: Vec<f32> = (0..self.config.dimensions)
95                .map(|_| normal.sample(&mut rng))
96                .collect();
97            self.relation_embeddings
98                .insert(relation.clone(), DVector::from_vec(embedding));
99        }
100
101        // Build adjacency matrix
102        self.build_adjacency_matrix(triples)?;
103
104        // Initialize weight matrices for each layer
105        self.weight_matrices.clear();
106        for _ in 0..self.num_layers {
107            let weight_matrix =
108                DMatrix::from_fn(self.config.dimensions, self.config.dimensions, |_, _| {
109                    normal.sample(&mut rng)
110                });
111            self.weight_matrices.push(weight_matrix);
112        }
113
114        Ok(())
115    }
116
117    /// Build adjacency matrix from triples
118    fn build_adjacency_matrix(&mut self, triples: &[Triple]) -> Result<()> {
119        let num_entities = self.entities.len();
120        let mut adj_matrix = DMatrix::zeros(num_entities, num_entities);
121
122        // Create entity index mapping
123        let entity_to_index: HashMap<String, usize> = self
124            .entities
125            .iter()
126            .enumerate()
127            .map(|(i, entity)| (entity.clone(), i))
128            .collect();
129
130        // Fill adjacency matrix
131        for triple in triples {
132            if let (Some(&subject_idx), Some(&object_idx)) = (
133                entity_to_index.get(&triple.subject),
134                entity_to_index.get(&triple.object),
135            ) {
136                adj_matrix[(subject_idx, object_idx)] = 1.0;
137                adj_matrix[(object_idx, subject_idx)] = 1.0; // Undirected graph
138            }
139        }
140
141        // Add self-loops
142        for i in 0..num_entities {
143            adj_matrix[(i, i)] = 1.0;
144        }
145
146        // Normalize adjacency matrix (symmetric normalization)
147        self.adjacency_matrix = Some(self.normalize_adjacency_matrix(adj_matrix));
148
149        Ok(())
150    }
151
152    /// Symmetric normalization of adjacency matrix: D^(-1/2) * A * D^(-1/2)
153    fn normalize_adjacency_matrix(&self, mut adj_matrix: DMatrix<f32>) -> DMatrix<f32> {
154        let num_nodes = adj_matrix.nrows();
155
156        // Calculate degree matrix
157        let mut degrees = Vec::with_capacity(num_nodes);
158        for i in 0..num_nodes {
159            let degree: f32 = (0..num_nodes).map(|j| adj_matrix[(i, j)]).sum();
160            degrees.push(if degree > 0.0 {
161                1.0 / degree.sqrt()
162            } else {
163                0.0
164            });
165        }
166
167        // Apply symmetric normalization
168        for i in 0..num_nodes {
169            for j in 0..num_nodes {
170                adj_matrix[(i, j)] *= degrees[i] * degrees[j];
171            }
172        }
173
174        adj_matrix
175    }
176
177    /// Forward pass through GCN layers
178    fn forward_pass(&self, features: &DMatrix<f32>) -> Result<DMatrix<f32>> {
179        let adj_matrix = self
180            .adjacency_matrix
181            .as_ref()
182            .ok_or_else(|| anyhow!("Adjacency matrix not initialized"))?;
183
184        let mut hidden = features.clone();
185
186        for layer_idx in 0..self.num_layers {
187            let weight = &self.weight_matrices[layer_idx];
188
189            // GCN layer: H^(l+1) = σ(A * H^(l) * W^(l))
190            let linear_transform = &hidden * weight;
191            hidden = adj_matrix * &linear_transform;
192
193            // Apply ReLU activation (except for last layer)
194            if layer_idx < self.num_layers - 1 {
195                hidden = hidden.map(|x| x.max(0.0));
196            }
197        }
198
199        Ok(hidden)
200    }
201
202    /// Train the GCN model
203    fn train_gcn(&mut self, _triples: &[Triple]) -> Result<()> {
204        // Create feature matrix from current embeddings
205        let num_entities = self.entities.len();
206        let mut features = DMatrix::zeros(num_entities, self.config.dimensions);
207
208        for (i, entity) in self.entities.iter().enumerate() {
209            if let Some(embedding) = self.entity_embeddings.get(entity) {
210                for (j, &value) in embedding.iter().enumerate() {
211                    features[(i, j)] = value;
212                }
213            }
214        }
215
216        // Perform forward pass
217        let updated_features = self.forward_pass(&features)?;
218
219        // Update entity embeddings with new features
220        for (i, entity) in self.entities.iter().enumerate() {
221            let new_embedding: Vec<f32> = (0..self.config.dimensions)
222                .map(|j| updated_features[(i, j)])
223                .collect();
224            self.entity_embeddings
225                .insert(entity.clone(), DVector::from_vec(new_embedding));
226        }
227
228        Ok(())
229    }
230}
231
232impl KGEmbeddingModel for GCN {
233    fn train(&mut self, triples: &[Triple]) -> Result<()> {
234        self.initialize(triples)?;
235
236        for epoch in 0..self.config.epochs {
237            self.train_gcn(triples)?;
238
239            if epoch % 10 == 0 {
240                println!("GCN training epoch {}/{}", epoch, self.config.epochs);
241            }
242        }
243
244        Ok(())
245    }
246
247    fn get_entity_embedding(&self, entity: &str) -> Option<Vector> {
248        self.entity_embeddings
249            .get(entity)
250            .map(|embedding| Vector::new(embedding.as_slice().to_vec()))
251    }
252
253    fn get_relation_embedding(&self, relation: &str) -> Option<Vector> {
254        self.relation_embeddings
255            .get(relation)
256            .map(|embedding| Vector::new(embedding.as_slice().to_vec()))
257    }
258
259    fn score_triple(&self, triple: &Triple) -> f32 {
260        // For GCN, we use cosine similarity between subject and object embeddings
261        // after considering the relation
262        if let (Some(subj_emb), Some(rel_emb), Some(obj_emb)) = (
263            self.get_entity_embedding(&triple.subject),
264            self.get_relation_embedding(&triple.predicate),
265            self.get_entity_embedding(&triple.object),
266        ) {
267            // Simple approach: h + r should be close to t
268            let predicted = subj_emb.add(&rel_emb).unwrap_or(subj_emb);
269            predicted.cosine_similarity(&obj_emb).unwrap_or(0.0)
270        } else {
271            0.0
272        }
273    }
274
275    fn predict_tail(&self, head: &str, relation: &str, k: usize) -> Vec<(String, f32)> {
276        if let (Some(head_emb), Some(rel_emb)) = (
277            self.get_entity_embedding(head),
278            self.get_relation_embedding(relation),
279        ) {
280            let query = head_emb.add(&rel_emb).unwrap_or(head_emb);
281
282            let mut scores = Vec::new();
283            for entity in &self.entities {
284                if entity != head {
285                    if let Some(entity_emb) = self.get_entity_embedding(entity) {
286                        let score = query.cosine_similarity(&entity_emb).unwrap_or(0.0);
287                        scores.push((entity.clone(), score));
288                    }
289                }
290            }
291
292            scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
293            scores.into_iter().take(k).collect()
294        } else {
295            Vec::new()
296        }
297    }
298
299    fn predict_head(&self, relation: &str, tail: &str, k: usize) -> Vec<(String, f32)> {
300        if let (Some(rel_emb), Some(tail_emb)) = (
301            self.get_relation_embedding(relation),
302            self.get_entity_embedding(tail),
303        ) {
304            let mut scores = Vec::new();
305            for entity in &self.entities {
306                if entity != tail {
307                    if let Some(entity_emb) = self.get_entity_embedding(entity) {
308                        let predicted = entity_emb.add(&rel_emb).unwrap_or(entity_emb);
309                        let score = predicted.cosine_similarity(&tail_emb).unwrap_or(0.0);
310                        scores.push((entity.clone(), score));
311                    }
312                }
313            }
314
315            scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
316            scores.into_iter().take(k).collect()
317        } else {
318            Vec::new()
319        }
320    }
321
322    fn get_entity_embeddings(&self) -> HashMap<String, Vector> {
323        self.entity_embeddings
324            .iter()
325            .map(|(entity, embedding)| (entity.clone(), Vector::new(embedding.as_slice().to_vec())))
326            .collect()
327    }
328
329    fn get_relation_embeddings(&self) -> HashMap<String, Vector> {
330        self.relation_embeddings
331            .iter()
332            .map(|(relation, embedding)| {
333                (relation.clone(), Vector::new(embedding.as_slice().to_vec()))
334            })
335            .collect()
336    }
337}
338
339/// GraphSAGE (Graph Sample and Aggregate) embedding model
340pub struct GraphSAGE {
341    config: KGEmbeddingConfig,
342    entity_embeddings: HashMap<String, DVector<f32>>,
343    relation_embeddings: HashMap<String, DVector<f32>>,
344    entities: Vec<String>,
345    relations: Vec<String>,
346    graph: HashMap<String, Vec<String>>, // Adjacency list
347    aggregator_type: AggregatorType,
348    num_layers: usize,
349    sample_size: usize,
350    sampling_strategy: SamplingStrategy,
351}
352
353#[derive(Debug, Clone, Copy)]
354pub enum AggregatorType {
355    Mean,
356    LSTM,
357    Pool,
358    Attention,
359}
360
361#[derive(Debug, Clone, Copy)]
362pub enum SamplingStrategy {
363    Uniform,  // Uniform random sampling
364    Degree,   // Degree-based sampling (prefer high-degree neighbors)
365    PageRank, // PageRank-based sampling (prefer important neighbors)
366    Recent,   // Sample recently added neighbors (for temporal graphs)
367}
368
369impl GraphSAGE {
370    pub fn new(config: KGEmbeddingConfig) -> Self {
371        Self {
372            config,
373            entity_embeddings: HashMap::new(),
374            relation_embeddings: HashMap::new(),
375            entities: Vec::new(),
376            relations: Vec::new(),
377            graph: HashMap::new(),
378            aggregator_type: AggregatorType::Mean,
379            num_layers: 2,
380            sample_size: 10, // Number of neighbors to sample
381            sampling_strategy: SamplingStrategy::Uniform,
382        }
383    }
384
385    pub fn with_aggregator(mut self, aggregator: AggregatorType) -> Self {
386        self.aggregator_type = aggregator;
387        self
388    }
389
390    pub fn with_sampling_strategy(mut self, strategy: SamplingStrategy) -> Self {
391        self.sampling_strategy = strategy;
392        self
393    }
394
395    pub fn with_sample_size(mut self, size: usize) -> Self {
396        self.sample_size = size;
397        self
398    }
399
400    /// Get embedding dimensions
401    pub fn dimensions(&self) -> usize {
402        self.config.dimensions
403    }
404
405    /// Initialize GraphSAGE model
406    fn initialize(&mut self, triples: &[Triple]) -> Result<()> {
407        // Collect unique entities and relations
408        let mut entities = std::collections::HashSet::new();
409        let mut relations = std::collections::HashSet::new();
410
411        for triple in triples {
412            entities.insert(triple.subject.clone());
413            entities.insert(triple.object.clone());
414            relations.insert(triple.predicate.clone());
415        }
416
417        self.entities = entities.into_iter().collect();
418        self.relations = relations.into_iter().collect();
419
420        // Build graph adjacency list
421        self.build_graph(triples);
422
423        // Initialize embeddings
424        let mut rng = if let Some(seed) = self.config.random_seed {
425            Random::seed(seed)
426        } else {
427            Random::seed(42)
428        };
429
430        let normal = Normal::new(0.0, 0.1)
431            .map_err(|e| anyhow!("Failed to create normal distribution: {}", e))?;
432
433        for entity in &self.entities {
434            let embedding: Vec<f32> = (0..self.config.dimensions)
435                .map(|_| normal.sample(&mut rng))
436                .collect();
437            self.entity_embeddings
438                .insert(entity.clone(), DVector::from_vec(embedding));
439        }
440
441        for relation in &self.relations {
442            let embedding: Vec<f32> = (0..self.config.dimensions)
443                .map(|_| normal.sample(&mut rng))
444                .collect();
445            self.relation_embeddings
446                .insert(relation.clone(), DVector::from_vec(embedding));
447        }
448
449        Ok(())
450    }
451
452    /// Build graph adjacency list
453    fn build_graph(&mut self, triples: &[Triple]) {
454        for triple in triples {
455            self.graph
456                .entry(triple.subject.clone())
457                .or_default()
458                .push(triple.object.clone());
459
460            self.graph
461                .entry(triple.object.clone())
462                .or_default()
463                .push(triple.subject.clone());
464        }
465    }
466
467    /// Sample neighbors for a node using different strategies
468    #[allow(deprecated)]
469    fn sample_neighbors(&self, node: &str, rng: &mut impl Rng) -> Vec<String> {
470        if let Some(neighbors) = self.graph.get(node) {
471            if neighbors.len() <= self.sample_size {
472                neighbors.clone()
473            } else {
474                match self.sampling_strategy {
475                    SamplingStrategy::Uniform => {
476                        // Note: Using manual random selection instead of SliceRandom
477                        // Manually sample neighbors using reservoir sampling
478                        let mut sampled = Vec::new();
479                        let sample_size = std::cmp::min(self.sample_size, neighbors.len());
480                        for (i, neighbor) in neighbors.iter().enumerate() {
481                            if sampled.len() < sample_size {
482                                sampled.push(neighbor.clone());
483                            } else {
484                                let j = rng.random_range(0..=i);
485                                if j < sample_size {
486                                    sampled[j] = neighbor.clone();
487                                }
488                            }
489                        }
490                        sampled
491                    }
492                    SamplingStrategy::Degree => self.degree_based_sampling(neighbors, rng),
493                    SamplingStrategy::PageRank => {
494                        // Simplified PageRank-based sampling (use degree as approximation)
495                        self.degree_based_sampling(neighbors, rng)
496                    }
497                    SamplingStrategy::Recent => {
498                        // For recent sampling, take the last added neighbors
499                        neighbors
500                            .iter()
501                            .rev()
502                            .take(self.sample_size)
503                            .cloned()
504                            .collect()
505                    }
506                }
507            }
508        } else {
509            Vec::new()
510        }
511    }
512
513    /// Degree-based sampling: prefer neighbors with higher degree
514    #[allow(deprecated)]
515    fn degree_based_sampling(&self, neighbors: &[String], rng: &mut impl Rng) -> Vec<String> {
516        let mut neighbor_degrees: Vec<(String, usize)> = neighbors
517            .iter()
518            .map(|neighbor| {
519                let degree = self.graph.get(neighbor).map(|n| n.len()).unwrap_or(0);
520                (neighbor.clone(), degree)
521            })
522            .collect();
523
524        // Sort by degree (descending) and add some randomization
525        neighbor_degrees.sort_by(|a, b| {
526            let degree_cmp = b.1.cmp(&a.1);
527            if degree_cmp == std::cmp::Ordering::Equal {
528                // Add randomization for ties
529                if rng.random_bool(0.5) {
530                    std::cmp::Ordering::Greater
531                } else {
532                    std::cmp::Ordering::Less
533                }
534            } else {
535                degree_cmp
536            }
537        });
538
539        neighbor_degrees
540            .into_iter()
541            .take(self.sample_size)
542            .map(|(neighbor, _)| neighbor)
543            .collect()
544    }
545
546    /// Aggregate neighbor embeddings
547    fn aggregate_neighbors(&self, neighbors: &[String]) -> Result<DVector<f32>> {
548        if neighbors.is_empty() {
549            return Ok(DVector::zeros(self.config.dimensions));
550        }
551
552        match self.aggregator_type {
553            AggregatorType::Mean => {
554                let mut sum = DVector::zeros(self.config.dimensions);
555                let mut count = 0;
556
557                for neighbor in neighbors {
558                    if let Some(embedding) = self.entity_embeddings.get(neighbor) {
559                        sum += embedding;
560                        count += 1;
561                    }
562                }
563
564                if count > 0 {
565                    Ok(sum / count as f32)
566                } else {
567                    Ok(DVector::zeros(self.config.dimensions))
568                }
569            }
570            AggregatorType::Pool => {
571                // Max pooling aggregator
572                let mut max_embedding =
573                    DVector::from_element(self.config.dimensions, f32::NEG_INFINITY);
574
575                for neighbor in neighbors {
576                    if let Some(embedding) = self.entity_embeddings.get(neighbor) {
577                        for i in 0..self.config.dimensions {
578                            max_embedding[i] = max_embedding[i].max(embedding[i]);
579                        }
580                    }
581                }
582
583                // Replace negative infinity with zeros
584                for i in 0..self.config.dimensions {
585                    if max_embedding[i] == f32::NEG_INFINITY {
586                        max_embedding[i] = 0.0;
587                    }
588                }
589
590                Ok(max_embedding)
591            }
592            AggregatorType::LSTM => {
593                // LSTM-based aggregator
594                self.lstm_aggregate(neighbors)
595            }
596            AggregatorType::Attention => {
597                // Attention-based aggregator
598                self.attention_aggregate(neighbors)
599            }
600        }
601    }
602
603    /// LSTM-based aggregator (simplified implementation)
604    fn lstm_aggregate(&self, neighbors: &[String]) -> Result<DVector<f32>> {
605        if neighbors.is_empty() {
606            return Ok(DVector::zeros(self.config.dimensions));
607        }
608
609        // Simplified LSTM: process neighbors sequentially with forget/input gates
610        let mut cell_state = DVector::zeros(self.config.dimensions);
611        let mut hidden_state = DVector::zeros(self.config.dimensions);
612
613        for neighbor in neighbors {
614            if let Some(embedding) = self.entity_embeddings.get(neighbor) {
615                // Simplified LSTM gates (using tanh and sigmoid approximations)
616                let forget_gate = embedding.map(|x| 1.0 / (1.0 + (-x).exp())); // sigmoid
617                let input_gate = embedding.map(|x| 1.0 / (1.0 + (-x).exp()));
618                let candidate = embedding.map(|x| x.tanh()); // tanh
619
620                // Update cell state
621                cell_state =
622                    cell_state.component_mul(&forget_gate) + input_gate.component_mul(&candidate);
623
624                // Update hidden state
625                let output_gate = embedding.map(|x| 1.0 / (1.0 + (-x).exp()));
626                hidden_state = output_gate.component_mul(&cell_state.map(|x| x.tanh()));
627            }
628        }
629
630        Ok(hidden_state)
631    }
632
633    /// Attention-based aggregator
634    fn attention_aggregate(&self, neighbors: &[String]) -> Result<DVector<f32>> {
635        if neighbors.is_empty() {
636            return Ok(DVector::zeros(self.config.dimensions));
637        }
638
639        let neighbor_embeddings: Vec<&DVector<f32>> = neighbors
640            .iter()
641            .filter_map(|neighbor| self.entity_embeddings.get(neighbor))
642            .collect();
643
644        if neighbor_embeddings.is_empty() {
645            return Ok(DVector::zeros(self.config.dimensions));
646        }
647
648        // Simple attention mechanism using dot-product attention
649        let mut attention_scores = Vec::new();
650        let mut weighted_sum = DVector::zeros(self.config.dimensions);
651
652        // Calculate attention scores (simplified: using magnitude as query)
653        let query = DVector::from_element(self.config.dimensions, 1.0); // Simple uniform query
654
655        for embedding in &neighbor_embeddings {
656            let score = query.dot(embedding).exp(); // Softmax will normalize
657            attention_scores.push(score);
658        }
659
660        // Normalize attention scores (softmax)
661        let total_score: f32 = attention_scores.iter().sum();
662        if total_score > 0.0 {
663            for score in &mut attention_scores {
664                *score /= total_score;
665            }
666        }
667
668        // Calculate weighted sum
669        for (embedding, &score) in neighbor_embeddings.iter().zip(attention_scores.iter()) {
670            weighted_sum += *embedding * score;
671        }
672
673        Ok(weighted_sum)
674    }
675
676    /// Forward pass for a single node
677    fn forward_node(&self, node: &str, rng: &mut impl Rng) -> Result<DVector<f32>> {
678        let neighbors = self.sample_neighbors(node, rng);
679        let neighbor_aggregate = self.aggregate_neighbors(&neighbors)?;
680
681        if let Some(node_embedding) = self.entity_embeddings.get(node) {
682            // Concatenate node embedding with aggregated neighbor embeddings
683            // For simplicity, we'll just add them (should be concatenation + linear transformation)
684            Ok(node_embedding + neighbor_aggregate)
685        } else {
686            Ok(neighbor_aggregate)
687        }
688    }
689}
690
691impl KGEmbeddingModel for GraphSAGE {
692    fn train(&mut self, triples: &[Triple]) -> Result<()> {
693        self.initialize(triples)?;
694
695        let mut rng = if let Some(seed) = self.config.random_seed {
696            Random::seed(seed)
697        } else {
698            Random::seed(42)
699        };
700
701        for epoch in 0..self.config.epochs {
702            let mut new_embeddings = HashMap::new();
703
704            // Update embeddings for all entities
705            for entity in &self.entities {
706                let new_embedding = self.forward_node(entity, &mut rng)?;
707                new_embeddings.insert(entity.clone(), new_embedding);
708            }
709
710            // Update embeddings
711            self.entity_embeddings = new_embeddings;
712
713            if epoch % 10 == 0 {
714                println!("GraphSAGE training epoch {}/{}", epoch, self.config.epochs);
715            }
716        }
717
718        Ok(())
719    }
720
721    fn get_entity_embedding(&self, entity: &str) -> Option<Vector> {
722        self.entity_embeddings
723            .get(entity)
724            .map(|embedding| Vector::new(embedding.as_slice().to_vec()))
725    }
726
727    fn get_relation_embedding(&self, relation: &str) -> Option<Vector> {
728        self.relation_embeddings
729            .get(relation)
730            .map(|embedding| Vector::new(embedding.as_slice().to_vec()))
731    }
732
733    fn score_triple(&self, triple: &Triple) -> f32 {
734        if let (Some(subj_emb), Some(rel_emb), Some(obj_emb)) = (
735            self.get_entity_embedding(&triple.subject),
736            self.get_relation_embedding(&triple.predicate),
737            self.get_entity_embedding(&triple.object),
738        ) {
739            let predicted = subj_emb.add(&rel_emb).unwrap_or(subj_emb);
740            predicted.cosine_similarity(&obj_emb).unwrap_or(0.0)
741        } else {
742            0.0
743        }
744    }
745
746    fn predict_tail(&self, head: &str, relation: &str, k: usize) -> Vec<(String, f32)> {
747        if let (Some(head_emb), Some(rel_emb)) = (
748            self.get_entity_embedding(head),
749            self.get_relation_embedding(relation),
750        ) {
751            let query = head_emb.add(&rel_emb).unwrap_or(head_emb);
752
753            let mut scores = Vec::new();
754            for entity in &self.entities {
755                if entity != head {
756                    if let Some(entity_emb) = self.get_entity_embedding(entity) {
757                        let score = query.cosine_similarity(&entity_emb).unwrap_or(0.0);
758                        scores.push((entity.clone(), score));
759                    }
760                }
761            }
762
763            scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
764            scores.into_iter().take(k).collect()
765        } else {
766            Vec::new()
767        }
768    }
769
770    fn predict_head(&self, relation: &str, tail: &str, k: usize) -> Vec<(String, f32)> {
771        if let (Some(rel_emb), Some(tail_emb)) = (
772            self.get_relation_embedding(relation),
773            self.get_entity_embedding(tail),
774        ) {
775            let mut scores = Vec::new();
776            for entity in &self.entities {
777                if entity != tail {
778                    if let Some(entity_emb) = self.get_entity_embedding(entity) {
779                        let predicted = entity_emb.add(&rel_emb).unwrap_or(entity_emb);
780                        let score = predicted.cosine_similarity(&tail_emb).unwrap_or(0.0);
781                        scores.push((entity.clone(), score));
782                    }
783                }
784            }
785
786            scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
787            scores.into_iter().take(k).collect()
788        } else {
789            Vec::new()
790        }
791    }
792
793    fn get_entity_embeddings(&self) -> HashMap<String, Vector> {
794        self.entity_embeddings
795            .iter()
796            .map(|(entity, embedding)| (entity.clone(), Vector::new(embedding.as_slice().to_vec())))
797            .collect()
798    }
799
800    fn get_relation_embeddings(&self) -> HashMap<String, Vector> {
801        self.relation_embeddings
802            .iter()
803            .map(|(relation, embedding)| {
804                (relation.clone(), Vector::new(embedding.as_slice().to_vec()))
805            })
806            .collect()
807    }
808}
809
810#[cfg(test)]
811mod tests {
812    use super::*;
813    use anyhow::Result;
814
815    #[test]
816    fn test_gcn_creation() {
817        let config = KGEmbeddingConfig {
818            model: crate::kg_embeddings::KGEmbeddingModelType::GCN,
819            dimensions: 64,
820            learning_rate: 0.01,
821            margin: 1.0,
822            negative_samples: 5,
823            batch_size: 32,
824            epochs: 10,
825            norm: 2,
826            random_seed: Some(42),
827            regularization: 0.01,
828        };
829
830        let gcn = GCN::new(config);
831        assert_eq!(gcn.num_layers, 2);
832    }
833
834    #[test]
835    fn test_graphsage_creation() {
836        let config = KGEmbeddingConfig {
837            model: crate::kg_embeddings::KGEmbeddingModelType::GraphSAGE,
838            dimensions: 64,
839            learning_rate: 0.01,
840            margin: 1.0,
841            negative_samples: 5,
842            batch_size: 32,
843            epochs: 10,
844            norm: 2,
845            random_seed: Some(42),
846            regularization: 0.01,
847        };
848
849        let graphsage = GraphSAGE::new(config);
850        assert_eq!(graphsage.sample_size, 10);
851    }
852
853    #[test]
854    fn test_gnn_training() -> Result<()> {
855        let config = KGEmbeddingConfig {
856            model: crate::kg_embeddings::KGEmbeddingModelType::GCN,
857            dimensions: 32,
858            learning_rate: 0.01,
859            margin: 1.0,
860            negative_samples: 5,
861            batch_size: 16,
862            epochs: 5,
863            norm: 2,
864            random_seed: Some(42),
865            regularization: 0.01,
866        };
867
868        let mut gcn = GCN::new(config);
869
870        let triples = vec![
871            Triple::new(
872                "entity1".to_string(),
873                "relation1".to_string(),
874                "entity2".to_string(),
875            ),
876            Triple::new(
877                "entity2".to_string(),
878                "relation2".to_string(),
879                "entity3".to_string(),
880            ),
881            Triple::new(
882                "entity1".to_string(),
883                "relation3".to_string(),
884                "entity3".to_string(),
885            ),
886        ];
887
888        // Should not panic
889        gcn.train(&triples)?;
890
891        // Should have embeddings for all entities
892        assert!(gcn.get_entity_embedding("entity1").is_some());
893        assert!(gcn.get_entity_embedding("entity2").is_some());
894        assert!(gcn.get_entity_embedding("entity3").is_some());
895        Ok(())
896    }
897}