Skip to main content

oxirs_embed/models/
gnn.rs

1//! Graph Neural Network (GNN) embedding models
2//!
3//! This module provides various GNN architectures for knowledge graph embeddings
4//! including GCN, GraphSAGE, GAT, and Graph Transformers.
5
6use crate::models::serialization::{MatrixF32, VectorF32};
7use crate::{
8    EmbeddingError, EmbeddingModel, ModelConfig, ModelStats, TrainingStats, Triple, Vector,
9};
10use anyhow::{anyhow, Result};
11use async_trait::async_trait;
12use chrono::{DateTime, Utc};
13use scirs2_core::ndarray_ext::{Array1, Array2};
14#[allow(unused_imports)]
15use scirs2_core::random::{Random, RngExt};
16use serde::{Deserialize, Serialize};
17use std::collections::{HashMap, HashSet};
18use std::fs::File;
19use std::io::{BufReader, BufWriter};
20use std::path::Path;
21use uuid::Uuid;
22
23/// Type of GNN architecture
24#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq)]
25pub enum GNNType {
26    /// Graph Convolutional Network
27    GCN,
28    /// GraphSAGE - Sampling and aggregating
29    GraphSAGE,
30    /// Graph Attention Network
31    GAT,
32    /// Graph Transformer
33    GraphTransformer,
34    /// Graph Isomorphism Network
35    GIN,
36    /// Principal Neighbourhood Aggregation
37    PNA,
38    /// Heterogeneous Graph Network
39    HetGNN,
40    /// Temporal Graph Network
41    TGN,
42}
43
44impl GNNType {
45    pub fn default_layers(&self) -> usize {
46        match self {
47            GNNType::GCN => 2,
48            GNNType::GraphSAGE => 2,
49            GNNType::GAT => 2,
50            GNNType::GraphTransformer => 4,
51            GNNType::GIN => 3,
52            GNNType::PNA => 3,
53            GNNType::HetGNN => 2,
54            GNNType::TGN => 2,
55        }
56    }
57
58    pub fn requires_attention(&self) -> bool {
59        matches!(self, GNNType::GAT | GNNType::GraphTransformer)
60    }
61}
62
63/// Aggregation method for GraphSAGE
64#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
65pub enum AggregationType {
66    Mean,
67    Max,
68    Sum,
69    LSTM,
70}
71
72/// Configuration for GNN models
73#[derive(Debug, Clone, Serialize, Deserialize)]
74pub struct GNNConfig {
75    pub base_config: ModelConfig,
76    pub gnn_type: GNNType,
77    pub num_layers: usize,
78    pub hidden_dimensions: Vec<usize>,
79    pub dropout: f64,
80    pub aggregation: AggregationType,
81    pub num_heads: Option<usize>,        // For attention-based models
82    pub sample_neighbors: Option<usize>, // For GraphSAGE
83    pub residual_connections: bool,
84    pub layer_norm: bool,
85    pub edge_features: bool,
86}
87
88impl Default for GNNConfig {
89    fn default() -> Self {
90        Self {
91            base_config: ModelConfig::default(),
92            gnn_type: GNNType::GCN,
93            num_layers: 2,
94            hidden_dimensions: vec![128, 64],
95            dropout: 0.1,
96            aggregation: AggregationType::Mean,
97            num_heads: None,
98            sample_neighbors: None,
99            residual_connections: true,
100            layer_norm: true,
101            edge_features: false,
102        }
103    }
104}
105
106/// GNN-based embedding model
107pub struct GNNEmbedding {
108    id: Uuid,
109    config: GNNConfig,
110    entity_embeddings: HashMap<String, Array1<f32>>,
111    relation_embeddings: HashMap<String, Array1<f32>>,
112    entity_to_idx: HashMap<String, usize>,
113    relation_to_idx: HashMap<String, usize>,
114    idx_to_entity: HashMap<usize, String>,
115    idx_to_relation: HashMap<usize, String>,
116    adjacency_list: HashMap<usize, HashSet<(usize, usize)>>, // (neighbor, relation)
117    reverse_adjacency_list: HashMap<usize, HashSet<(usize, usize)>>,
118    triples: Vec<Triple>,
119    layers: Vec<GNNLayer>,
120    is_trained: bool,
121    creation_time: chrono::DateTime<Utc>,
122    last_training_time: Option<chrono::DateTime<Utc>>,
123}
124
125/// Single GNN layer
126struct GNNLayer {
127    weight_matrix: Array2<f32>,
128    bias: Array1<f32>,
129    attention_weights: Option<AttentionWeights>,
130    layer_norm: Option<LayerNormalization>,
131}
132
133/// Attention weights for GAT/GraphTransformer
134struct AttentionWeights {
135    query_weights: Array2<f32>,
136    key_weights: Array2<f32>,
137    value_weights: Array2<f32>,
138    num_heads: usize,
139}
140
141/// Layer normalization parameters
142struct LayerNormalization {
143    gamma: Array1<f32>,
144    beta: Array1<f32>,
145    epsilon: f32,
146}
147
148/// Serializable mirror of [`AttentionWeights`].
149#[derive(Debug, Serialize, Deserialize)]
150struct AttentionWeightsSer {
151    query_weights: MatrixF32,
152    key_weights: MatrixF32,
153    value_weights: MatrixF32,
154    num_heads: usize,
155}
156
157/// Serializable mirror of [`LayerNormalization`].
158#[derive(Debug, Serialize, Deserialize)]
159struct LayerNormalizationSer {
160    gamma: VectorF32,
161    beta: VectorF32,
162    epsilon: f32,
163}
164
165/// Serializable mirror of a [`GNNLayer`].
166#[derive(Debug, Serialize, Deserialize)]
167struct GNNLayerSer {
168    weight_matrix: MatrixF32,
169    bias: VectorF32,
170    attention_weights: Option<AttentionWeightsSer>,
171    layer_norm: Option<LayerNormalizationSer>,
172}
173
174impl GNNLayerSer {
175    fn from_layer(layer: &GNNLayer) -> Self {
176        Self {
177            weight_matrix: MatrixF32::from_array(&layer.weight_matrix),
178            bias: VectorF32::from_array(&layer.bias),
179            attention_weights: layer
180                .attention_weights
181                .as_ref()
182                .map(|a| AttentionWeightsSer {
183                    query_weights: MatrixF32::from_array(&a.query_weights),
184                    key_weights: MatrixF32::from_array(&a.key_weights),
185                    value_weights: MatrixF32::from_array(&a.value_weights),
186                    num_heads: a.num_heads,
187                }),
188            layer_norm: layer.layer_norm.as_ref().map(|l| LayerNormalizationSer {
189                gamma: VectorF32::from_array(&l.gamma),
190                beta: VectorF32::from_array(&l.beta),
191                epsilon: l.epsilon,
192            }),
193        }
194    }
195
196    fn into_layer(self) -> Result<GNNLayer> {
197        let attention_weights = match self.attention_weights {
198            Some(a) => Some(AttentionWeights {
199                query_weights: a.query_weights.to_array()?,
200                key_weights: a.key_weights.to_array()?,
201                value_weights: a.value_weights.to_array()?,
202                num_heads: a.num_heads,
203            }),
204            None => None,
205        };
206        let layer_norm = self.layer_norm.map(|l| LayerNormalization {
207            gamma: l.gamma.to_array(),
208            beta: l.beta.to_array(),
209            epsilon: l.epsilon,
210        });
211        Ok(GNNLayer {
212            weight_matrix: self.weight_matrix.to_array()?,
213            bias: self.bias.to_array(),
214            attention_weights,
215            layer_norm,
216        })
217    }
218}
219
220/// Serializable representation of a [`GNNEmbedding`] model for persistence.
221#[derive(Debug, Serialize, Deserialize)]
222struct GNNSerializable {
223    id: Uuid,
224    config: GNNConfig,
225    entity_embeddings: HashMap<String, Vec<f32>>,
226    relation_embeddings: HashMap<String, Vec<f32>>,
227    entity_to_idx: HashMap<String, usize>,
228    relation_to_idx: HashMap<String, usize>,
229    idx_to_entity: HashMap<usize, String>,
230    idx_to_relation: HashMap<usize, String>,
231    adjacency_list: HashMap<usize, HashSet<(usize, usize)>>,
232    reverse_adjacency_list: HashMap<usize, HashSet<(usize, usize)>>,
233    triples: Vec<Triple>,
234    layers: Vec<GNNLayerSer>,
235    is_trained: bool,
236    creation_time: DateTime<Utc>,
237    last_training_time: Option<DateTime<Utc>>,
238}
239
240impl GNNEmbedding {
241    pub fn new(config: GNNConfig) -> Self {
242        Self {
243            id: Uuid::new_v4(),
244            config,
245            entity_embeddings: HashMap::new(),
246            relation_embeddings: HashMap::new(),
247            entity_to_idx: HashMap::new(),
248            relation_to_idx: HashMap::new(),
249            idx_to_entity: HashMap::new(),
250            idx_to_relation: HashMap::new(),
251            adjacency_list: HashMap::new(),
252            reverse_adjacency_list: HashMap::new(),
253            triples: Vec::new(),
254            layers: Vec::new(),
255            is_trained: false,
256            creation_time: Utc::now(),
257            last_training_time: None,
258        }
259    }
260
261    /// Initialize GNN layers
262    fn initialize_layers(&mut self) -> Result<()> {
263        self.layers.clear();
264        let mut rng = Random::seed(42);
265
266        let mut input_dim = self.config.base_config.dimensions;
267        let num_layers = self.config.num_layers;
268
269        for i in 0..num_layers {
270            let output_dim = if i == num_layers - 1 {
271                // Final layer should output back to original embedding dimension
272                self.config.base_config.dimensions
273            } else if i < self.config.hidden_dimensions.len() {
274                self.config.hidden_dimensions[i]
275            } else {
276                self.config.base_config.dimensions
277            };
278
279            // Initialize weight matrix
280            let scale = (2.0 / (input_dim + output_dim) as f32).sqrt();
281            let weight_matrix = Array2::from_shape_fn((input_dim, output_dim), |_| {
282                rng.random_range(0.0..1.0) * scale * 2.0 - scale
283            });
284
285            let bias = Array1::zeros(output_dim);
286
287            // Initialize attention weights if needed
288            let attention_weights = if self.config.gnn_type.requires_attention() {
289                let num_heads = self.config.num_heads.unwrap_or(8);
290                let head_dim = output_dim / num_heads;
291
292                // For multi-head attention, each head processes a portion of the output
293                let attention_dim = head_dim * num_heads; // Should equal output_dim
294
295                Some(AttentionWeights {
296                    query_weights: Array2::from_shape_fn((input_dim, attention_dim), |_| {
297                        rng.random_range(0.0..1.0) * scale * 2.0 - scale
298                    }),
299                    key_weights: Array2::from_shape_fn((input_dim, attention_dim), |_| {
300                        rng.random_range(0.0..1.0) * scale * 2.0 - scale
301                    }),
302                    value_weights: Array2::from_shape_fn((input_dim, attention_dim), |_| {
303                        rng.random_range(0.0..1.0) * scale * 2.0 - scale
304                    }),
305                    num_heads,
306                })
307            } else {
308                None
309            };
310
311            // Initialize layer normalization if needed
312            let layer_norm = if self.config.layer_norm {
313                Some(LayerNormalization {
314                    gamma: Array1::ones(output_dim),
315                    beta: Array1::zeros(output_dim),
316                    epsilon: 1e-5,
317                })
318            } else {
319                None
320            };
321
322            self.layers.push(GNNLayer {
323                weight_matrix,
324                bias,
325                attention_weights,
326                layer_norm,
327            });
328
329            input_dim = output_dim;
330        }
331
332        Ok(())
333    }
334
335    /// Build adjacency lists from triples
336    fn build_adjacency_lists(&mut self) {
337        self.adjacency_list.clear();
338        self.reverse_adjacency_list.clear();
339
340        for triple in &self.triples {
341            let subject_idx = self.entity_to_idx[&triple.subject.iri];
342            let object_idx = self.entity_to_idx[&triple.object.iri];
343            let relation_idx = self.relation_to_idx[&triple.predicate.iri];
344
345            // Forward adjacency
346            self.adjacency_list
347                .entry(subject_idx)
348                .or_default()
349                .insert((object_idx, relation_idx));
350
351            // Reverse adjacency
352            self.reverse_adjacency_list
353                .entry(object_idx)
354                .or_default()
355                .insert((subject_idx, relation_idx));
356        }
357    }
358
359    /// Aggregate neighbor features
360    fn aggregate_neighbors(
361        &self,
362        node_idx: usize,
363        node_features: &HashMap<usize, Array1<f32>>,
364    ) -> Array1<f32> {
365        let neighbors = self.adjacency_list.get(&node_idx);
366        let reverse_neighbors = self.reverse_adjacency_list.get(&node_idx);
367
368        let mut neighbor_features = Vec::new();
369
370        // Collect forward neighbors
371        if let Some(neighbors) = neighbors {
372            for (neighbor_idx, _) in neighbors {
373                if let Some(feature) = node_features.get(neighbor_idx) {
374                    neighbor_features.push(feature.clone());
375                }
376            }
377        }
378
379        // Collect reverse neighbors
380        if let Some(reverse_neighbors) = reverse_neighbors {
381            for (neighbor_idx, _) in reverse_neighbors {
382                if let Some(feature) = node_features.get(neighbor_idx) {
383                    neighbor_features.push(feature.clone());
384                }
385            }
386        }
387
388        if neighbor_features.is_empty() {
389            // Return zero vector if no neighbors
390            return Array1::zeros(
391                node_features
392                    .values()
393                    .next()
394                    .expect("node_features should not be empty")
395                    .len(),
396            );
397        }
398
399        // Aggregate based on configuration
400        match self.config.aggregation {
401            AggregationType::Mean => {
402                let sum: Array1<f32> = neighbor_features
403                    .iter()
404                    .fold(Array1::zeros(neighbor_features[0].len()), |acc, x| acc + x);
405                sum / neighbor_features.len() as f32
406            }
407            AggregationType::Max => neighbor_features.iter().fold(
408                Array1::from_elem(neighbor_features[0].len(), f32::NEG_INFINITY),
409                |acc, x| {
410                    let mut result = acc.clone();
411                    for (i, &val) in x.iter().enumerate() {
412                        result[i] = result[i].max(val);
413                    }
414                    result
415                },
416            ),
417            AggregationType::Sum => neighbor_features
418                .iter()
419                .fold(Array1::zeros(neighbor_features[0].len()), |acc, x| acc + x),
420            AggregationType::LSTM => {
421                // Simplified LSTM aggregation - in practice would use actual LSTM
422                self.aggregate_neighbors_lstm(&neighbor_features)
423            }
424        }
425    }
426
427    /// LSTM aggregation (simplified)
428    fn aggregate_neighbors_lstm(&self, neighbor_features: &[Array1<f32>]) -> Array1<f32> {
429        // Simplified version - real implementation would use LSTM cells
430        let mut aggregated = Array1::zeros(neighbor_features[0].len());
431        for feature in neighbor_features {
432            aggregated = aggregated * 0.8 + feature * 0.2; // Simple weighted average
433        }
434        aggregated
435    }
436
437    /// Apply GNN layer
438    fn apply_layer(
439        &self,
440        layer: &GNNLayer,
441        node_features: &HashMap<usize, Array1<f32>>,
442    ) -> HashMap<usize, Array1<f32>> {
443        let mut new_features = HashMap::new();
444
445        match self.config.gnn_type {
446            GNNType::GCN => self.apply_gcn_layer(layer, node_features, &mut new_features),
447            GNNType::GraphSAGE => {
448                self.apply_graphsage_layer(layer, node_features, &mut new_features)
449            }
450            GNNType::GAT => self.apply_gat_layer(layer, node_features, &mut new_features),
451            GNNType::GIN => self.apply_gin_layer(layer, node_features, &mut new_features),
452            _ => self.apply_gcn_layer(layer, node_features, &mut new_features), // Default to GCN
453        }
454
455        new_features
456    }
457
458    /// Apply GCN layer
459    fn apply_gcn_layer(
460        &self,
461        layer: &GNNLayer,
462        node_features: &HashMap<usize, Array1<f32>>,
463        new_features: &mut HashMap<usize, Array1<f32>>,
464    ) {
465        for (node_idx, feature) in node_features {
466            let aggregated = self.aggregate_neighbors(*node_idx, node_features);
467            let combined = feature + &aggregated;
468            let transformed = combined.dot(&layer.weight_matrix) + &layer.bias;
469
470            // Apply activation (ReLU)
471            let activated = transformed.mapv(|x| x.max(0.0));
472
473            // Apply layer norm if configured
474            let output = if let Some(ln) = &layer.layer_norm {
475                self.apply_layer_norm(&activated, ln)
476            } else {
477                activated
478            };
479
480            new_features.insert(*node_idx, output);
481        }
482    }
483
484    /// Apply GraphSAGE layer
485    fn apply_graphsage_layer(
486        &self,
487        layer: &GNNLayer,
488        node_features: &HashMap<usize, Array1<f32>>,
489        new_features: &mut HashMap<usize, Array1<f32>>,
490    ) {
491        for (node_idx, feature) in node_features {
492            let aggregated = self.aggregate_neighbors(*node_idx, node_features);
493
494            // For GraphSAGE, we apply separate transformations and then combine
495            // Transform node feature
496            let node_transformed = feature.dot(&layer.weight_matrix) + &layer.bias;
497
498            // Transform aggregated neighbor features (reuse same weight matrix for simplicity)
499            let neighbor_transformed = aggregated.dot(&layer.weight_matrix) + &layer.bias;
500
501            // Combine the transformed features
502            let combined = &node_transformed + &neighbor_transformed;
503
504            // Apply activation and normalization
505            let activated = combined.mapv(|x| x.max(0.0));
506            let normalized = &activated / (activated.dot(&activated).sqrt() + 1e-6);
507
508            new_features.insert(*node_idx, normalized);
509        }
510    }
511
512    /// Apply GAT layer
513    fn apply_gat_layer(
514        &self,
515        layer: &GNNLayer,
516        node_features: &HashMap<usize, Array1<f32>>,
517        new_features: &mut HashMap<usize, Array1<f32>>,
518    ) {
519        // Simplified GAT - real implementation would compute attention scores
520        let attention = layer
521            .attention_weights
522            .as_ref()
523            .expect("attention_weights should be initialized for GAT layer");
524
525        for (node_idx, feature) in node_features {
526            // Get neighbors
527            let mut neighbor_indices = Vec::new();
528            if let Some(neighbors) = self.adjacency_list.get(node_idx) {
529                neighbor_indices.extend(neighbors.iter().map(|(n, _)| *n));
530            }
531            if let Some(neighbors) = self.reverse_adjacency_list.get(node_idx) {
532                neighbor_indices.extend(neighbors.iter().map(|(n, _)| *n));
533            }
534
535            if neighbor_indices.is_empty() {
536                // Apply linear transformation even when no neighbors
537                let transformed = feature.dot(&layer.weight_matrix) + &layer.bias;
538                let activated = transformed.mapv(|x| x.max(0.0));
539                new_features.insert(*node_idx, activated);
540                continue;
541            }
542
543            // Ensure feature dimensions match weight matrix input dimensions
544            if feature.len() != attention.query_weights.shape()[0] {
545                // Fallback to simple aggregation if dimensions don't match
546                let aggregated = self.aggregate_neighbors(*node_idx, node_features);
547                let combined = feature + &aggregated;
548                let transformed = combined.dot(&layer.weight_matrix) + &layer.bias;
549                let activated = transformed.mapv(|x| x.max(0.0));
550                new_features.insert(*node_idx, activated);
551                continue;
552            }
553
554            // Compute attention scores (simplified)
555            let query = feature.dot(&attention.query_weights);
556            let mut attention_scores = Vec::new();
557            let mut neighbor_values = Vec::new();
558
559            for neighbor_idx in &neighbor_indices {
560                if let Some(neighbor_feature) = node_features.get(neighbor_idx) {
561                    // Check dimension compatibility before computing attention
562                    if neighbor_feature.len() != attention.key_weights.shape()[0] {
563                        continue;
564                    }
565
566                    let key = neighbor_feature.dot(&attention.key_weights);
567                    let value = neighbor_feature.dot(&attention.value_weights);
568
569                    // Compute attention score with proper dimension checking
570                    if query.len() == key.len() {
571                        let score = query.dot(&key) / (attention.num_heads as f32).sqrt();
572                        attention_scores.push(score);
573                        neighbor_values.push(value);
574                    }
575                }
576            }
577
578            if attention_scores.is_empty() {
579                // Fallback to simple aggregation if no valid attention scores
580                let aggregated = self.aggregate_neighbors(*node_idx, node_features);
581                let combined = feature + &aggregated;
582                let transformed = combined.dot(&layer.weight_matrix) + &layer.bias;
583                let activated = transformed.mapv(|x| x.max(0.0));
584                new_features.insert(*node_idx, activated);
585                continue;
586            }
587
588            // Softmax
589            let max_score = attention_scores
590                .iter()
591                .fold(f32::NEG_INFINITY, |a, &b| a.max(b));
592            let exp_scores: Vec<f32> = attention_scores
593                .iter()
594                .map(|&s| (s - max_score).exp())
595                .collect();
596            let sum_exp = exp_scores.iter().sum::<f32>();
597            let attention_weights: Vec<f32> =
598                exp_scores.iter().copied().map(|e| e / sum_exp).collect();
599
600            // Apply attention with proper output dimensions
601            let output_dim = layer.weight_matrix.shape()[1];
602            let mut aggregated = Array1::<f32>::zeros(output_dim);
603
604            for (i, value) in neighbor_values.iter().enumerate() {
605                // Ensure value dimension matches output dimension
606                let min_dim = aggregated.len().min(value.len());
607                for j in 0..min_dim {
608                    aggregated[j] += value[j] * attention_weights[i];
609                }
610            }
611
612            // Apply linear transformation
613            let transformed = feature.dot(&layer.weight_matrix) + &layer.bias;
614            let combined =
615                if self.config.residual_connections && transformed.len() == aggregated.len() {
616                    transformed + &aggregated
617                } else {
618                    transformed
619                };
620
621            let activated = combined.mapv(|x| x.max(0.0));
622            new_features.insert(*node_idx, activated);
623        }
624    }
625
626    /// Apply GIN layer
627    fn apply_gin_layer(
628        &self,
629        layer: &GNNLayer,
630        node_features: &HashMap<usize, Array1<f32>>,
631        new_features: &mut HashMap<usize, Array1<f32>>,
632    ) {
633        let epsilon = 0.0; // GIN epsilon parameter
634
635        for (node_idx, feature) in node_features {
636            let aggregated = self.aggregate_neighbors(*node_idx, node_features);
637            let combined = (1.0 + epsilon) * feature + aggregated;
638
639            // MLP transformation (simplified as single linear layer)
640            let transformed = combined.dot(&layer.weight_matrix) + &layer.bias;
641            let activated = transformed.mapv(|x| x.max(0.0));
642
643            new_features.insert(*node_idx, activated);
644        }
645    }
646
647    /// Apply layer normalization
648    fn apply_layer_norm(&self, input: &Array1<f32>, ln: &LayerNormalization) -> Array1<f32> {
649        let mean = input.mean().unwrap_or(0.0);
650        let variance = input.mapv(|x| (x - mean).powi(2)).mean().unwrap_or(1.0);
651        let normalized = input.mapv(|x| (x - mean) / (variance + ln.epsilon).sqrt());
652        &normalized * &ln.gamma + &ln.beta
653    }
654
655    /// Forward pass through all GNN layers
656    fn forward(
657        &self,
658        initial_features: HashMap<usize, Array1<f32>>,
659    ) -> HashMap<usize, Array1<f32>> {
660        let mut features = initial_features;
661
662        for layer in self.layers.iter() {
663            let new_features = self.apply_layer(layer, &features);
664
665            // Apply dropout during training (simplified - always applied here)
666            let dropout_rate = self.config.dropout;
667            let mut rng = Random::seed(42);
668
669            features = new_features
670                .into_iter()
671                .map(|(idx, feat)| {
672                    let masked = feat.mapv(|x| {
673                        if rng.random_range(0.0..1.0) > dropout_rate as f32 {
674                            x / (1.0 - dropout_rate as f32)
675                        } else {
676                            0.0
677                        }
678                    });
679                    (idx, masked)
680                })
681                .collect();
682        }
683
684        features
685    }
686}
687
688#[async_trait]
689impl EmbeddingModel for GNNEmbedding {
690    fn config(&self) -> &ModelConfig {
691        &self.config.base_config
692    }
693
694    fn model_id(&self) -> &Uuid {
695        &self.id
696    }
697
698    fn model_type(&self) -> &'static str {
699        "GNNEmbedding"
700    }
701
702    fn add_triple(&mut self, triple: Triple) -> Result<()> {
703        // Add entities to index
704        let subject = triple.subject.iri.clone();
705        let object = triple.object.iri.clone();
706        let predicate = triple.predicate.iri.clone();
707
708        if !self.entity_to_idx.contains_key(&subject) {
709            let idx = self.entity_to_idx.len();
710            self.entity_to_idx.insert(subject.clone(), idx);
711            self.idx_to_entity.insert(idx, subject);
712        }
713
714        if !self.entity_to_idx.contains_key(&object) {
715            let idx = self.entity_to_idx.len();
716            self.entity_to_idx.insert(object.clone(), idx);
717            self.idx_to_entity.insert(idx, object);
718        }
719
720        if !self.relation_to_idx.contains_key(&predicate) {
721            let idx = self.relation_to_idx.len();
722            self.relation_to_idx.insert(predicate.clone(), idx);
723            self.idx_to_relation.insert(idx, predicate);
724        }
725
726        self.triples.push(triple);
727        self.is_trained = false;
728        Ok(())
729    }
730
731    async fn train(&mut self, epochs: Option<usize>) -> Result<TrainingStats> {
732        let start_time = std::time::Instant::now();
733        let epochs = epochs.unwrap_or(self.config.base_config.max_epochs);
734
735        // Build adjacency lists
736        self.build_adjacency_lists();
737
738        // Initialize layers
739        self.initialize_layers()?;
740
741        // Initialize random embeddings
742        let mut rng = Random::seed(42);
743        let dimensions = self.config.base_config.dimensions;
744
745        let mut initial_features = HashMap::new();
746        for idx in self.entity_to_idx.values() {
747            let embedding =
748                Array1::from_shape_fn(dimensions, |_| rng.random_range(0.0..1.0) * 0.1 - 0.05);
749            initial_features.insert(*idx, embedding);
750        }
751
752        // Training loop (simplified)
753        let mut loss_history = Vec::new();
754
755        for _epoch in 0..epochs {
756            // Forward pass
757            let output_features = self.forward(initial_features.clone());
758
759            // Compute loss (simplified - just using L2 regularization)
760            let loss = output_features
761                .values()
762                .map(|f| f.mapv(|x| x * x).sum())
763                .sum::<f32>()
764                / output_features.len() as f32;
765
766            loss_history.push(loss as f64);
767
768            // Update initial features with output (simplified training)
769            initial_features = output_features;
770
771            // Early stopping
772            if loss < 0.001 {
773                break;
774            }
775        }
776
777        // Store final embeddings
778        for (idx, embedding) in initial_features {
779            if let Some(entity) = self.idx_to_entity.get(&idx) {
780                self.entity_embeddings.insert(entity.clone(), embedding);
781            }
782        }
783
784        // Generate relation embeddings (simplified - using random initialization)
785        for relation in self.relation_to_idx.keys() {
786            let embedding =
787                Array1::from_shape_fn(dimensions, |_| rng.random_range(0.0..1.0) * 0.1 - 0.05);
788            self.relation_embeddings.insert(relation.clone(), embedding);
789        }
790
791        self.is_trained = true;
792        self.last_training_time = Some(Utc::now());
793
794        Ok(TrainingStats {
795            epochs_completed: loss_history.len(),
796            final_loss: *loss_history.last().unwrap_or(&0.0),
797            training_time_seconds: start_time.elapsed().as_secs_f64(),
798            convergence_achieved: loss_history.last().unwrap_or(&1.0) < &0.001,
799            loss_history,
800        })
801    }
802
803    fn get_entity_embedding(&self, entity: &str) -> Result<Vector> {
804        if !self.is_trained {
805            return Err(EmbeddingError::ModelNotTrained.into());
806        }
807
808        self.entity_embeddings
809            .get(entity)
810            .map(|e| Vector::new(e.to_vec()))
811            .ok_or_else(|| {
812                EmbeddingError::EntityNotFound {
813                    entity: entity.to_string(),
814                }
815                .into()
816            })
817    }
818
819    fn get_relation_embedding(&self, relation: &str) -> Result<Vector> {
820        if !self.is_trained {
821            return Err(EmbeddingError::ModelNotTrained.into());
822        }
823
824        self.relation_embeddings
825            .get(relation)
826            .map(|e| Vector::new(e.to_vec()))
827            .ok_or_else(|| {
828                EmbeddingError::RelationNotFound {
829                    relation: relation.to_string(),
830                }
831                .into()
832            })
833    }
834
835    fn score_triple(&self, subject: &str, predicate: &str, object: &str) -> Result<f64> {
836        if !self.is_trained {
837            return Err(EmbeddingError::ModelNotTrained.into());
838        }
839
840        let subj_emb =
841            self.entity_embeddings
842                .get(subject)
843                .ok_or_else(|| EmbeddingError::EntityNotFound {
844                    entity: subject.to_string(),
845                })?;
846
847        let pred_emb = self.relation_embeddings.get(predicate).ok_or_else(|| {
848            EmbeddingError::RelationNotFound {
849                relation: predicate.to_string(),
850            }
851        })?;
852
853        let obj_emb =
854            self.entity_embeddings
855                .get(object)
856                .ok_or_else(|| EmbeddingError::EntityNotFound {
857                    entity: object.to_string(),
858                })?;
859
860        // Simple scoring: dot product of transformed embeddings
861        let transformed = (subj_emb + pred_emb) * obj_emb;
862        Ok(transformed.sum() as f64)
863    }
864
865    fn predict_objects(
866        &self,
867        subject: &str,
868        predicate: &str,
869        k: usize,
870    ) -> Result<Vec<(String, f64)>> {
871        if !self.is_trained {
872            return Err(EmbeddingError::ModelNotTrained.into());
873        }
874
875        let mut scores = Vec::new();
876
877        for entity in self.entity_to_idx.keys() {
878            if let Ok(score) = self.score_triple(subject, predicate, entity) {
879                scores.push((entity.clone(), score));
880            }
881        }
882
883        scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
884        scores.truncate(k);
885
886        Ok(scores)
887    }
888
889    fn predict_subjects(
890        &self,
891        predicate: &str,
892        object: &str,
893        k: usize,
894    ) -> Result<Vec<(String, f64)>> {
895        if !self.is_trained {
896            return Err(EmbeddingError::ModelNotTrained.into());
897        }
898
899        let mut scores = Vec::new();
900
901        for entity in self.entity_to_idx.keys() {
902            if let Ok(score) = self.score_triple(entity, predicate, object) {
903                scores.push((entity.clone(), score));
904            }
905        }
906
907        scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
908        scores.truncate(k);
909
910        Ok(scores)
911    }
912
913    fn predict_relations(
914        &self,
915        subject: &str,
916        object: &str,
917        k: usize,
918    ) -> Result<Vec<(String, f64)>> {
919        if !self.is_trained {
920            return Err(EmbeddingError::ModelNotTrained.into());
921        }
922
923        let mut scores = Vec::new();
924
925        for relation in self.relation_to_idx.keys() {
926            if let Ok(score) = self.score_triple(subject, relation, object) {
927                scores.push((relation.clone(), score));
928            }
929        }
930
931        scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
932        scores.truncate(k);
933
934        Ok(scores)
935    }
936
937    fn get_entities(&self) -> Vec<String> {
938        self.entity_to_idx.keys().cloned().collect()
939    }
940
941    fn get_relations(&self) -> Vec<String> {
942        self.relation_to_idx.keys().cloned().collect()
943    }
944
945    fn get_stats(&self) -> ModelStats {
946        ModelStats {
947            num_entities: self.entity_to_idx.len(),
948            num_relations: self.relation_to_idx.len(),
949            num_triples: self.triples.len(),
950            dimensions: self.config.base_config.dimensions,
951            is_trained: self.is_trained,
952            model_type: format!("GNNEmbedding-{:?}", self.config.gnn_type),
953            creation_time: self.creation_time,
954            last_training_time: self.last_training_time,
955        }
956    }
957
958    fn save(&self, path: &str) -> Result<()> {
959        let serializable = GNNSerializable {
960            id: self.id,
961            config: self.config.clone(),
962            entity_embeddings: self
963                .entity_embeddings
964                .iter()
965                .map(|(k, v)| (k.clone(), v.to_vec()))
966                .collect(),
967            relation_embeddings: self
968                .relation_embeddings
969                .iter()
970                .map(|(k, v)| (k.clone(), v.to_vec()))
971                .collect(),
972            entity_to_idx: self.entity_to_idx.clone(),
973            relation_to_idx: self.relation_to_idx.clone(),
974            idx_to_entity: self.idx_to_entity.clone(),
975            idx_to_relation: self.idx_to_relation.clone(),
976            adjacency_list: self.adjacency_list.clone(),
977            reverse_adjacency_list: self.reverse_adjacency_list.clone(),
978            triples: self.triples.clone(),
979            layers: self.layers.iter().map(GNNLayerSer::from_layer).collect(),
980            is_trained: self.is_trained,
981            creation_time: self.creation_time,
982            last_training_time: self.last_training_time,
983        };
984
985        let file = File::create(path)
986            .map_err(|e| anyhow!("Failed to create model file {}: {}", path, e))?;
987        let writer = BufWriter::new(file);
988        oxicode::serde::encode_into_std_write(&serializable, writer, oxicode::config::standard())
989            .map_err(|e| anyhow!("Failed to serialize GNN model: {}", e))?;
990        Ok(())
991    }
992
993    fn load(&mut self, path: &str) -> Result<()> {
994        if !Path::new(path).exists() {
995            return Err(anyhow!("Model file not found: {}", path));
996        }
997
998        let file =
999            File::open(path).map_err(|e| anyhow!("Failed to open model file {}: {}", path, e))?;
1000        let reader = BufReader::new(file);
1001        let (serializable, _): (GNNSerializable, _) =
1002            oxicode::serde::decode_from_std_read(reader, oxicode::config::standard())
1003                .map_err(|e| anyhow!("Failed to deserialize GNN model: {}", e))?;
1004
1005        self.id = serializable.id;
1006        self.config = serializable.config;
1007        self.entity_embeddings = serializable
1008            .entity_embeddings
1009            .into_iter()
1010            .map(|(k, v)| (k, Array1::from_vec(v)))
1011            .collect();
1012        self.relation_embeddings = serializable
1013            .relation_embeddings
1014            .into_iter()
1015            .map(|(k, v)| (k, Array1::from_vec(v)))
1016            .collect();
1017        self.entity_to_idx = serializable.entity_to_idx;
1018        self.relation_to_idx = serializable.relation_to_idx;
1019        self.idx_to_entity = serializable.idx_to_entity;
1020        self.idx_to_relation = serializable.idx_to_relation;
1021        self.adjacency_list = serializable.adjacency_list;
1022        self.reverse_adjacency_list = serializable.reverse_adjacency_list;
1023        self.triples = serializable.triples;
1024        self.layers = serializable
1025            .layers
1026            .into_iter()
1027            .map(GNNLayerSer::into_layer)
1028            .collect::<Result<Vec<_>>>()?;
1029        self.is_trained = serializable.is_trained;
1030        self.creation_time = serializable.creation_time;
1031        self.last_training_time = serializable.last_training_time;
1032        Ok(())
1033    }
1034
1035    fn clear(&mut self) {
1036        self.entity_embeddings.clear();
1037        self.relation_embeddings.clear();
1038        self.entity_to_idx.clear();
1039        self.relation_to_idx.clear();
1040        self.idx_to_entity.clear();
1041        self.idx_to_relation.clear();
1042        self.adjacency_list.clear();
1043        self.reverse_adjacency_list.clear();
1044        self.triples.clear();
1045        self.layers.clear();
1046        self.is_trained = false;
1047    }
1048
1049    fn is_trained(&self) -> bool {
1050        self.is_trained
1051    }
1052
1053    async fn encode(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>> {
1054        Err(anyhow!(
1055            "Knowledge graph embedding model does not support text encoding"
1056        ))
1057    }
1058}
1059
1060#[cfg(test)]
1061mod tests {
1062    use super::*;
1063    use crate::NamedNode;
1064
1065    #[tokio::test]
1066    async fn test_gnn_embedding_basic() {
1067        let config = GNNConfig {
1068            gnn_type: GNNType::GCN,
1069            num_layers: 2,
1070            hidden_dimensions: vec![64, 32],
1071            ..Default::default()
1072        };
1073
1074        let mut model = GNNEmbedding::new(config);
1075
1076        // Add some triples
1077        let triple1 = Triple::new(
1078            NamedNode::new("http://example.org/Alice").expect("should succeed"),
1079            NamedNode::new("http://example.org/knows").expect("should succeed"),
1080            NamedNode::new("http://example.org/Bob").expect("should succeed"),
1081        );
1082
1083        let triple2 = Triple::new(
1084            NamedNode::new("http://example.org/Bob").expect("should succeed"),
1085            NamedNode::new("http://example.org/knows").expect("should succeed"),
1086            NamedNode::new("http://example.org/Charlie").expect("should succeed"),
1087        );
1088
1089        model.add_triple(triple1).expect("should succeed");
1090        model.add_triple(triple2).expect("should succeed");
1091
1092        // Train the model
1093        let _stats = model.train(Some(10)).await.expect("should succeed");
1094        assert!(model.is_trained());
1095
1096        // Get embeddings
1097        let alice_emb = model
1098            .get_entity_embedding("http://example.org/Alice")
1099            .expect("should succeed");
1100        assert_eq!(alice_emb.dimensions, 100); // Default dimensions
1101
1102        // Test predictions
1103        let predictions = model
1104            .predict_objects("http://example.org/Alice", "http://example.org/knows", 5)
1105            .expect("should succeed");
1106        assert!(!predictions.is_empty());
1107    }
1108
1109    #[tokio::test]
1110    async fn test_gnn_types() {
1111        for gnn_type in [GNNType::GCN, GNNType::GraphSAGE, GNNType::GAT, GNNType::GIN] {
1112            let config = GNNConfig {
1113                gnn_type,
1114                num_heads: if gnn_type == GNNType::GAT {
1115                    Some(4)
1116                } else {
1117                    None
1118                },
1119                ..Default::default()
1120            };
1121
1122            let mut model = GNNEmbedding::new(config);
1123
1124            let triple = Triple::new(
1125                NamedNode::new("http://example.org/A").expect("should succeed"),
1126                NamedNode::new("http://example.org/rel").expect("should succeed"),
1127                NamedNode::new("http://example.org/B").expect("should succeed"),
1128            );
1129
1130            model.add_triple(triple).expect("should succeed");
1131            let _stats = model.train(Some(5)).await.expect("should succeed");
1132            assert!(model.is_trained());
1133        }
1134    }
1135}