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

144 lines
3.6 KiB
Rust

use rtx_polygraph::{
error::PolygraphError,
ir::{DataType, IRNode, IRNodeType, NodeId, Shape},
};
#[test]
fn test_ir_node_creation() {
let node = IRNode::new(
NodeId(0),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(1), NodeId(2)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
assert_eq!(node.id(), NodeId(0));
assert_eq!(node.inputs().len(), 2);
assert_eq!(node.outputs().len(), 1);
assert!(matches!(node.node_type(), IRNodeType::MatMul { .. }));
}
#[test]
fn test_ir_node_sparse_operation() {
let node = IRNode::new(
NodeId(10),
IRNodeType::SparseMatMul {
format: rtx_polygraph::ir::SparseFormat::CSR,
},
vec![NodeId(11), NodeId(12)],
vec![DataType::F32],
vec![Shape::new(vec![256, 512])],
);
assert_eq!(node.id(), NodeId(10));
assert!(matches!(node.node_type(), IRNodeType::SparseMatMul { .. }));
}
#[test]
fn test_ir_node_graph_operation() {
let node = IRNode::new(
NodeId(20),
IRNodeType::GraphConvolution {
aggregation: rtx_polygraph::ir::AggregationType::Sum,
},
vec![NodeId(21), NodeId(22)],
vec![DataType::F32],
vec![Shape::new(vec![100, 64])],
);
assert_eq!(node.id(), NodeId(20));
assert!(matches!(
node.node_type(),
IRNodeType::GraphConvolution { .. }
));
}
#[test]
fn test_ir_node_fft_operation() {
let node = IRNode::new(
NodeId(30),
IRNodeType::FFT {
inverse: false,
axes: vec![0, 1],
},
vec![NodeId(31)],
vec![DataType::C64],
vec![Shape::new(vec![128, 128])],
);
assert_eq!(node.id(), NodeId(30));
assert!(matches!(node.node_type(), IRNodeType::FFT { .. }));
}
#[test]
fn test_ir_node_control_flow() {
let node = IRNode::new(
NodeId(40),
IRNodeType::Conditional {
condition_id: NodeId(41),
},
vec![NodeId(42), NodeId(43)],
vec![DataType::F32],
vec![Shape::new(vec![64, 64])],
);
assert_eq!(node.id(), NodeId(40));
assert!(matches!(node.node_type(), IRNodeType::Conditional { .. }));
}
#[test]
fn test_node_signature_generation() {
let node = IRNode::new(
NodeId(0),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(1), NodeId(2)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
let signature = node.signature();
assert!(!signature.is_empty());
// Same configuration should produce same signature
let node2 = IRNode::new(
NodeId(100), // Different ID should not affect signature
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![NodeId(101), NodeId(102)],
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
assert_eq!(node.signature(), node2.signature());
}
#[test]
fn test_invalid_node_creation_fails() {
// Empty inputs for operation that requires inputs should fail validation
let result = IRNode::try_new(
NodeId(0),
IRNodeType::MatMul {
transpose_a: false,
transpose_b: true,
},
vec![], // Empty inputs - should fail
vec![DataType::F32],
vec![Shape::new(vec![64, 128])],
);
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
PolygraphError::InvalidNodeConfiguration { .. }
));
}