use approx::assert_relative_eq; use rtx_geom::layers::{GATLayer, GCNLayer, GraphSAGELayer}; use rtx_geom::{Edge, GNNLayer, Graph, Node}; use rtx_tensor::{Device, Tensor}; #[test] fn test_gcn_layer_forward() { let mut graph = Graph::new(); let device = Device::cpu(); // Create simple graph with 3 nodes, 2D features let node0 = graph.add_node(Node::new( Tensor::from_data(vec![1.0, 0.0], [2], &device).unwrap(), )); let node1 = graph.add_node(Node::new( Tensor::from_data(vec![0.0, 1.0], [2], &device).unwrap(), )); let node2 = graph.add_node(Node::new( Tensor::from_data(vec![1.0, 1.0], [2], &device).unwrap(), )); // Connect them graph .add_edge( node0, node1, Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); graph .add_edge( node1, node2, Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); graph .add_edge( node0, node2, Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); // Create GCN layer: 2 input features -> 4 output features let mut gcn_layer = GCNLayer::new(2, 4).unwrap(); let output = gcn_layer.forward(&graph).unwrap(); assert_eq!(output.len(), 3); // 3 nodes for node_output in &output { assert_eq!(node_output.shape().dims(), &[4]); // 4 output features per node } } #[test] #[ignore = "Pre-existing device mismatch (CPU vs Metal) error"] fn test_gat_layer_attention() { let mut graph = Graph::new(); let device = Device::cpu(); let node0 = graph.add_node(Node::new( Tensor::from_data(vec![1.0, 2.0, 3.0], [3], &device).unwrap(), )); let node1 = graph.add_node(Node::new( Tensor::from_data(vec![4.0, 5.0, 6.0], [3], &device).unwrap(), )); let node2 = graph.add_node(Node::new( Tensor::from_data(vec![7.0, 8.0, 9.0], [3], &device).unwrap(), )); graph .add_edge( node0, node1, Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); graph .add_edge( node1, node2, Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); // GAT layer: 3 input -> 8 output, 4 attention heads (8 % 4 = 2 features per head) let mut gat_layer = GATLayer::new(3, 8, 4).unwrap(); let output = gat_layer.forward(&graph).unwrap(); assert_eq!(output.len(), 3); // 3 nodes for node_output in &output { assert_eq!(node_output.shape().dims(), &[8]); // 8 output features per node } // Test that attention weights are computed let attention_weights = gat_layer.get_attention_weights(); assert!(!attention_weights.is_empty()); } #[test] #[ignore = "Pre-existing device mismatch (CPU vs Metal) error"] fn test_graphsage_sampling() { let mut graph = Graph::new(); let device = Device::cpu(); // Create larger graph for sampling let mut nodes = Vec::new(); for i in 0..10 { let features = Tensor::from_data(vec![i as f32, (i * 2) as f32], [2], &device).unwrap(); nodes.push(graph.add_node(Node::new(features))); } // Connect nodes in a chain for i in 0..9 { graph .add_edge( nodes[i], nodes[i + 1], Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); } // GraphSAGE layer with sampling let mut sage_layer = GraphSAGELayer::new(2, 3, 2).unwrap(); // sample_size = 2 let output = sage_layer.forward(&graph).unwrap(); assert_eq!(output.len(), 10); // 10 nodes for node_output in &output { assert_eq!(node_output.shape().dims(), &[3]); // 3 output features per node } } #[test] fn test_layer_parameter_initialization() { // Test that layers properly initialize their parameters let gcn = GCNLayer::new(5, 10).unwrap(); let params = gcn.parameters(); // Should have weight matrix W and bias vector b assert_eq!(params.len(), 2); assert_eq!(params[0].shape().dims(), &[5, 10]); // Weight matrix assert_eq!(params[1].shape().dims(), &[10]); // Bias vector // Parameters should be initialized (not all zeros) let weight_data = params[0].to_cpu().unwrap(); let non_zero_count = weight_data.iter().filter(|&&x| x.abs() > 1e-8).count(); assert!( non_zero_count > 0, "Parameters should be initialized with non-zero values" ); } #[test] fn test_layer_gradient_flow() { let mut graph = Graph::new(); let device = Device::cpu(); let node0 = graph.add_node(Node::new( Tensor::from_data(vec![1.0, 2.0], [2], &device).unwrap(), )); let node1 = graph.add_node(Node::new( Tensor::from_data(vec![3.0, 4.0], [2], &device).unwrap(), )); graph .add_edge( node0, node1, Edge::new(Tensor::from_data(vec![1.0], [1], &device).unwrap()), ) .unwrap(); let mut gcn_layer = GCNLayer::new(2, 3).unwrap(); gcn_layer.enable_gradients(true); let output = gcn_layer.forward(&graph).unwrap(); // Simulate backward pass with dummy gradients let mut grad_outputs = Vec::new(); for _ in &output { grad_outputs.push(Tensor::ones([3], &device).unwrap()); } let _param_grads = gcn_layer.backward(&grad_outputs).unwrap(); // Test that gradients were computed for parameters assert!(gcn_layer.has_computed_gradients()); }