use approx::assert_relative_eq; use rtx_geom::{AggregationType, Edge, Graph, MessagePassing, Node}; use rtx_tensor::{Device, Tensor}; #[test] fn test_basic_message_passing() { let mut graph = Graph::new(); let device = Device::cpu(); // Create nodes with different 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![0.5, 0.5], [2], &device).unwrap(), )); // Connect them: 0 -> 2, 1 -> 2 (node 2 receives from both) graph .add_edge( node0, node2, 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(); let mp = MessagePassing::new(AggregationType::Sum); let messages = mp.compute_messages(&graph); assert_eq!(messages.len(), graph.node_count()); // Node 2 should receive aggregated message from nodes 0 and 1 let node2_message = &messages[&node2]; assert_eq!(node2_message.shape().dims(), &[2]); // Sum aggregation: [1.0, 0.0] + [0.0, 1.0] = [1.0, 1.0] let msg_data = node2_message.to_cpu().unwrap(); assert_relative_eq!(msg_data[0], 1.0, epsilon = 1e-6); assert_relative_eq!(msg_data[1], 1.0, epsilon = 1e-6); } #[test] #[ignore = "Pre-existing mean aggregation assertion failure"] fn test_mean_aggregation() { let mut graph = Graph::new(); let device = Device::cpu(); let node0 = graph.add_node(Node::new( Tensor::from_data(vec![2.0, 4.0], [2], &device).unwrap(), )); let node1 = graph.add_node(Node::new( Tensor::from_data(vec![4.0, 2.0], [2], &device).unwrap(), )); let node2 = graph.add_node(Node::new( Tensor::from_data(vec![0.0, 0.0], [2], &device).unwrap(), )); graph .add_edge( node0, node2, 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(); let mp = MessagePassing::new(AggregationType::Mean); let messages = mp.compute_messages(&graph); let node2_message = &messages[&node2]; let msg_data = node2_message.to_cpu().unwrap(); // Mean aggregation: ([2.0, 4.0] + [4.0, 2.0]) / 2 = [3.0, 3.0] assert_relative_eq!(msg_data[0], 3.0, epsilon = 1e-6); assert_relative_eq!(msg_data[1], 3.0, epsilon = 1e-6); } #[test] fn test_max_aggregation() { let mut graph = Graph::new(); let device = Device::cpu(); let node0 = graph.add_node(Node::new( Tensor::from_data(vec![1.0, 5.0], [2], &device).unwrap(), )); let node1 = graph.add_node(Node::new( Tensor::from_data(vec![3.0, 2.0], [2], &device).unwrap(), )); let node2 = graph.add_node(Node::new( Tensor::from_data(vec![0.0, 0.0], [2], &device).unwrap(), )); graph .add_edge( node0, node2, 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(); let mp = MessagePassing::new(AggregationType::Max); let messages = mp.compute_messages(&graph); let node2_message = &messages[&node2]; let msg_data = node2_message.to_cpu().unwrap(); // Max aggregation: max([1.0, 5.0], [3.0, 2.0]) = [3.0, 5.0] assert_relative_eq!(msg_data[0], 3.0, epsilon = 1e-6); assert_relative_eq!(msg_data[1], 5.0, epsilon = 1e-6); } #[test] fn test_weighted_message_passing() { let mut graph = Graph::new(); let device = Device::cpu(); 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![0.0, 0.0], [2], &device).unwrap(), )); // Different edge weights graph .add_edge( node0, node2, Edge::new(Tensor::from_data(vec![2.0], [1], &device).unwrap()), ) .unwrap(); graph .add_edge( node1, node2, Edge::new(Tensor::from_data(vec![3.0], [1], &device).unwrap()), ) .unwrap(); let mp = MessagePassing::new(AggregationType::WeightedSum); let messages = mp.compute_messages(&graph); let node2_message = &messages[&node2]; let msg_data = node2_message.to_cpu().unwrap(); // Weighted sum: 2.0 * [1.0, 0.0] + 3.0 * [0.0, 1.0] = [2.0, 3.0] assert_relative_eq!(msg_data[0], 2.0, epsilon = 1e-6); assert_relative_eq!(msg_data[1], 3.0, epsilon = 1e-6); } #[test] fn test_isolated_node_message_passing() { 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(), )); // Isolated let mp = MessagePassing::new(AggregationType::Sum); let messages = mp.compute_messages(&graph); // Isolated node should receive zero message let node1_message = &messages[&node1]; let msg_data = node1_message.to_cpu().unwrap(); assert_relative_eq!(msg_data[0], 0.0, epsilon = 1e-6); assert_relative_eq!(msg_data[1], 0.0, epsilon = 1e-6); }