1use 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#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq)]
25pub enum GNNType {
26 GCN,
28 GraphSAGE,
30 GAT,
32 GraphTransformer,
34 GIN,
36 PNA,
38 HetGNN,
40 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#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
65pub enum AggregationType {
66 Mean,
67 Max,
68 Sum,
69 LSTM,
70}
71
72#[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>, pub sample_neighbors: Option<usize>, 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
106pub 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)>>, 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
125struct GNNLayer {
127 weight_matrix: Array2<f32>,
128 bias: Array1<f32>,
129 attention_weights: Option<AttentionWeights>,
130 layer_norm: Option<LayerNormalization>,
131}
132
133struct AttentionWeights {
135 query_weights: Array2<f32>,
136 key_weights: Array2<f32>,
137 value_weights: Array2<f32>,
138 num_heads: usize,
139}
140
141struct LayerNormalization {
143 gamma: Array1<f32>,
144 beta: Array1<f32>,
145 epsilon: f32,
146}
147
148#[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#[derive(Debug, Serialize, Deserialize)]
159struct LayerNormalizationSer {
160 gamma: VectorF32,
161 beta: VectorF32,
162 epsilon: f32,
163}
164
165#[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#[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 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 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 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 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 let attention_dim = head_dim * num_heads; 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 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 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 self.adjacency_list
347 .entry(subject_idx)
348 .or_default()
349 .insert((object_idx, relation_idx));
350
351 self.reverse_adjacency_list
353 .entry(object_idx)
354 .or_default()
355 .insert((subject_idx, relation_idx));
356 }
357 }
358
359 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 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 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 Array1::zeros(
391 node_features
392 .values()
393 .next()
394 .expect("node_features should not be empty")
395 .len(),
396 );
397 }
398
399 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 self.aggregate_neighbors_lstm(&neighbor_features)
423 }
424 }
425 }
426
427 fn aggregate_neighbors_lstm(&self, neighbor_features: &[Array1<f32>]) -> Array1<f32> {
429 let mut aggregated = Array1::zeros(neighbor_features[0].len());
431 for feature in neighbor_features {
432 aggregated = aggregated * 0.8 + feature * 0.2; }
434 aggregated
435 }
436
437 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), }
454
455 new_features
456 }
457
458 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 let activated = transformed.mapv(|x| x.max(0.0));
472
473 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 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 let node_transformed = feature.dot(&layer.weight_matrix) + &layer.bias;
497
498 let neighbor_transformed = aggregated.dot(&layer.weight_matrix) + &layer.bias;
500
501 let combined = &node_transformed + &neighbor_transformed;
503
504 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 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 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 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 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 if feature.len() != attention.query_weights.shape()[0] {
545 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 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 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 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 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 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 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 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 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 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; 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 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 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 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 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 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 self.build_adjacency_lists();
737
738 self.initialize_layers()?;
740
741 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 let mut loss_history = Vec::new();
754
755 for _epoch in 0..epochs {
756 let output_features = self.forward(initial_features.clone());
758
759 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 initial_features = output_features;
770
771 if loss < 0.001 {
773 break;
774 }
775 }
776
777 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 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 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 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 let _stats = model.train(Some(10)).await.expect("should succeed");
1094 assert!(model.is_trained());
1095
1096 let alice_emb = model
1098 .get_entity_embedding("http://example.org/Alice")
1099 .expect("should succeed");
1100 assert_eq!(alice_emb.dimensions, 100); 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}