144 lines
3.6 KiB
Rust
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 { .. }
|
|
));
|
|
}
|