Files
rustytorch/crates/specialized/rtx-geom/tests/message_passing.rs
T
2026-03-04 00:08:42 +00:00

199 lines
5.9 KiB
Rust

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);
}