1use 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
17pub 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; 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 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 fn initialize(&mut self, triples: &[Triple]) -> Result<()> {
60 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 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 self.build_adjacency_matrix(triples)?;
103
104 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 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 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 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; }
139 }
140
141 for i in 0..num_entities {
143 adj_matrix[(i, i)] = 1.0;
144 }
145
146 self.adjacency_matrix = Some(self.normalize_adjacency_matrix(adj_matrix));
148
149 Ok(())
150 }
151
152 fn normalize_adjacency_matrix(&self, mut adj_matrix: DMatrix<f32>) -> DMatrix<f32> {
154 let num_nodes = adj_matrix.nrows();
155
156 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 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 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 let linear_transform = &hidden * weight;
191 hidden = adj_matrix * &linear_transform;
192
193 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 fn train_gcn(&mut self, _triples: &[Triple]) -> Result<()> {
204 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 let updated_features = self.forward_pass(&features)?;
218
219 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 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 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
339pub 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>>, 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, Degree, PageRank, Recent, }
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, 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 pub fn dimensions(&self) -> usize {
402 self.config.dimensions
403 }
404
405 fn initialize(&mut self, triples: &[Triple]) -> Result<()> {
407 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 self.build_graph(triples);
422
423 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 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 #[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 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 self.degree_based_sampling(neighbors, rng)
496 }
497 SamplingStrategy::Recent => {
498 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 #[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 neighbor_degrees.sort_by(|a, b| {
526 let degree_cmp = b.1.cmp(&a.1);
527 if degree_cmp == std::cmp::Ordering::Equal {
528 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 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 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 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 self.lstm_aggregate(neighbors)
595 }
596 AggregatorType::Attention => {
597 self.attention_aggregate(neighbors)
599 }
600 }
601 }
602
603 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 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 let forget_gate = embedding.map(|x| 1.0 / (1.0 + (-x).exp())); let input_gate = embedding.map(|x| 1.0 / (1.0 + (-x).exp()));
618 let candidate = embedding.map(|x| x.tanh()); cell_state =
622 cell_state.component_mul(&forget_gate) + input_gate.component_mul(&candidate);
623
624 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 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 let mut attention_scores = Vec::new();
650 let mut weighted_sum = DVector::zeros(self.config.dimensions);
651
652 let query = DVector::from_element(self.config.dimensions, 1.0); for embedding in &neighbor_embeddings {
656 let score = query.dot(embedding).exp(); attention_scores.push(score);
658 }
659
660 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 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 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 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 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 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 gcn.train(&triples)?;
890
891 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}