#![allow(dead_code)]
use crate::parameter::Parameter;
use crate::{GraphData, GraphLayer};
use torsh_tensor::{
creation::{randn, zeros},
Tensor,
};
use scirs2_core::ndarray::{Array1, Array2, ArrayView1, Axis};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Debug)]
pub struct MPNNConv {
in_features: usize,
out_features: usize,
edge_features: usize,
message_hidden_dim: usize,
update_hidden_dim: usize,
message_layer1: Parameter,
message_layer2: Parameter,
message_bias1: Option<Parameter>,
message_bias2: Option<Parameter>,
update_layer1: Parameter,
update_layer2: Parameter,
update_bias1: Option<Parameter>,
update_bias2: Option<Parameter>,
edge_embedding: Option<Parameter>,
aggregation_type: AggregationType,
}
#[derive(Debug, Clone, Copy)]
pub enum AggregationType {
Sum,
Mean,
Max,
Attention,
}
impl MPNNConv {
pub fn new(
in_features: usize,
out_features: usize,
edge_features: usize,
message_hidden_dim: usize,
update_hidden_dim: usize,
aggregation_type: AggregationType,
bias: bool,
) -> Self {
let message_input_dim = 2 * in_features + edge_features;
let message_layer1 = Parameter::new(
randn(&[message_input_dim, message_hidden_dim])
.expect("failed to create message layer 1 weights"),
);
let message_layer2 = Parameter::new(
randn(&[message_hidden_dim, out_features])
.expect("failed to create message layer 2 weights"),
);
let message_bias1 = if bias {
Some(Parameter::new(
zeros(&[message_hidden_dim]).expect("failed to create message bias 1"),
))
} else {
None
};
let message_bias2 = if bias {
Some(Parameter::new(
zeros(&[out_features]).expect("failed to create message bias 2"),
))
} else {
None
};
let update_input_dim = in_features + out_features;
let update_layer1 = Parameter::new(
randn(&[update_input_dim, update_hidden_dim])
.expect("failed to create update layer 1 weights"),
);
let update_layer2 = Parameter::new(
randn(&[update_hidden_dim, out_features])
.expect("failed to create update layer 2 weights"),
);
let update_bias1 = if bias {
Some(Parameter::new(
zeros(&[update_hidden_dim]).expect("failed to create update bias 1"),
))
} else {
None
};
let update_bias2 = if bias {
Some(Parameter::new(
zeros(&[out_features]).expect("failed to create update bias 2"),
))
} else {
None
};
let edge_embedding = if edge_features > 0 {
Some(Parameter::new(
randn(&[edge_features, edge_features])
.expect("failed to create edge embedding weights"),
))
} else {
None
};
Self {
in_features,
out_features,
edge_features,
message_hidden_dim,
update_hidden_dim,
message_layer1,
message_layer2,
message_bias1,
message_bias2,
update_layer1,
update_layer2,
update_bias1,
update_bias2,
edge_embedding,
aggregation_type,
}
}
pub fn forward(&self, graph: &GraphData) -> GraphData {
let num_nodes = graph.num_nodes;
let edge_data = crate::utils::tensor_to_vec2::<f32>(&graph.edge_index)
.expect("failed to extract edge index data");
let _num_edges = edge_data[0].len();
let messages = self.compute_messages(graph);
let aggregated = self.aggregate_messages(&messages, &edge_data, num_nodes);
let updated_features = self.update_nodes(&graph.x, &aggregated);
GraphData {
x: updated_features,
edge_index: graph.edge_index.clone(),
edge_attr: graph.edge_attr.clone(),
batch: graph.batch.clone(),
num_nodes: graph.num_nodes,
num_edges: graph.num_edges,
}
}
fn compute_messages(&self, graph: &GraphData) -> Tensor {
let edge_data = crate::utils::tensor_to_vec2::<f32>(&graph.edge_index)
.expect("failed to extract edge index data");
let num_edges = edge_data[0].len();
let mut all_messages = Vec::new();
for edge_idx in 0..num_edges {
let src_idx = edge_data[0][edge_idx] as usize;
let dst_idx = edge_data[1][edge_idx] as usize;
let h_i = graph
.x
.slice_tensor(0, src_idx, src_idx + 1)
.expect("failed to slice source node features")
.squeeze_tensor(0)
.expect("failed to squeeze source node features");
let h_j = graph
.x
.slice_tensor(0, dst_idx, dst_idx + 1)
.expect("failed to slice destination node features")
.squeeze_tensor(0)
.expect("failed to squeeze destination node features");
let edge_feat = if let Some(ref edge_attr) = graph.edge_attr {
if self.edge_features > 0 {
let e_ij = edge_attr
.slice_tensor(0, edge_idx, edge_idx + 1)
.expect("failed to slice edge attributes")
.squeeze_tensor(0)
.expect("failed to squeeze edge attributes");
if let Some(ref edge_emb) = self.edge_embedding {
let e_ij_2d = e_ij
.unsqueeze_tensor(0)
.expect("failed to unsqueeze edge features");
e_ij_2d
.matmul(&edge_emb.clone_data())
.expect("failed to apply edge embedding")
.squeeze_tensor(0)
.expect("failed to squeeze embedded edge features")
} else {
e_ij
}
} else {
zeros(&[self.edge_features]).expect("failed to create zero edge features")
}
} else {
zeros(&[self.edge_features]).expect("failed to create zero edge features")
};
let message_input = Tensor::cat(&[&h_i, &h_j, &edge_feat], 0)
.expect("failed to concatenate message input");
let message_input_2d = message_input
.unsqueeze_tensor(0)
.expect("failed to unsqueeze message input");
let mut message = message_input_2d
.matmul(&self.message_layer1.clone_data())
.expect("failed to apply message layer 1")
.squeeze_tensor(0)
.expect("failed to squeeze message layer 1 output");
if let Some(ref bias1) = self.message_bias1 {
message = message
.add(&bias1.clone_data())
.expect("operation should succeed");
}
message = message
.maximum(
&zeros(&message.shape().dims()).expect("failed to create zero tensor for ReLU"),
)
.expect("failed to apply ReLU activation");
let message_2d = message
.unsqueeze_tensor(0)
.expect("failed to unsqueeze message for layer 2");
message = message_2d
.matmul(&self.message_layer2.clone_data())
.expect("failed to apply message layer 2")
.squeeze_tensor(0)
.expect("failed to squeeze message layer 2 output");
if let Some(ref bias2) = self.message_bias2 {
message = message
.add(&bias2.clone_data())
.expect("operation should succeed");
}
all_messages.push(message);
}
if all_messages.is_empty() {
zeros(&[0, self.out_features]).expect("failed to create empty messages tensor")
} else {
let mut message_data = Vec::new();
for msg in &all_messages {
let msg_vec = msg.to_vec().expect("conversion should succeed");
message_data.extend(msg_vec);
}
torsh_tensor::creation::from_vec(
message_data,
&[all_messages.len(), self.out_features],
torsh_core::device::DeviceType::Cpu,
)
.expect("failed to create messages tensor from data")
}
}
fn aggregate_messages(
&self,
messages: &Tensor,
edge_data: &[Vec<f32>],
num_nodes: usize,
) -> Tensor {
let mut aggregated = zeros(&[num_nodes, self.out_features])
.expect("failed to create aggregated messages tensor");
let num_edges = edge_data[0].len();
if num_edges == 0 {
return aggregated;
}
match self.aggregation_type {
AggregationType::Sum | AggregationType::Mean => {
let mut node_counts = vec![0; num_nodes];
for edge_idx in 0..num_edges {
let dst_idx = edge_data[1][edge_idx] as usize;
if dst_idx < num_nodes {
let message = messages
.slice_tensor(0, edge_idx, edge_idx + 1)
.expect("failed to slice message")
.squeeze_tensor(0)
.expect("failed to squeeze message");
let current = aggregated
.slice_tensor(0, dst_idx, dst_idx + 1)
.expect("failed to slice aggregated tensor")
.squeeze_tensor(0)
.expect("failed to squeeze aggregated tensor");
let updated = current.add(&message).expect("operation should succeed");
aggregated
.slice_tensor(0, dst_idx, dst_idx + 1)
.expect("failed to slice aggregated tensor for update")
.copy_(
&updated
.unsqueeze_tensor(0)
.expect("failed to unsqueeze updated tensor"),
)
.expect("failed to copy updated tensor");
node_counts[dst_idx] += 1;
}
}
if matches!(self.aggregation_type, AggregationType::Mean) {
for node in 0..num_nodes {
if node_counts[node] > 0 {
let current = aggregated
.slice_tensor(0, node, node + 1)
.expect("failed to slice aggregated tensor for mean")
.squeeze_tensor(0)
.expect("failed to squeeze aggregated tensor for mean");
let normalized = current
.div_scalar(node_counts[node] as f32)
.expect("failed to normalize aggregated tensor");
aggregated
.slice_tensor(0, node, node + 1)
.expect("failed to slice aggregated tensor for normalized update")
.copy_(
&normalized
.unsqueeze_tensor(0)
.expect("failed to unsqueeze normalized tensor"),
)
.expect("failed to copy normalized tensor");
}
}
}
}
AggregationType::Max => {
aggregated
.fill_(-1e9_f32)
.expect("failed to fill aggregated tensor with initial values");
for edge_idx in 0..num_edges {
let dst_idx = edge_data[1][edge_idx] as usize;
if dst_idx < num_nodes {
let message = messages
.slice_tensor(0, edge_idx, edge_idx + 1)
.expect("failed to slice message for max aggregation")
.squeeze_tensor(0)
.expect("failed to squeeze message for max aggregation");
let current = aggregated
.slice_tensor(0, dst_idx, dst_idx + 1)
.expect("failed to slice aggregated tensor for max")
.squeeze_tensor(0)
.expect("failed to squeeze aggregated tensor for max");
let updated = current
.maximum(&message)
.expect("failed to compute maximum");
aggregated
.slice_tensor(0, dst_idx, dst_idx + 1)
.expect("failed to slice aggregated tensor for max update")
.copy_(
&updated
.unsqueeze_tensor(0)
.expect("failed to unsqueeze max updated tensor"),
)
.expect("failed to copy max updated tensor");
}
}
let aggregated_data = aggregated.to_vec().expect("conversion should succeed");
let filtered_data: Vec<f32> = aggregated_data
.iter()
.map(|&x| if x <= -1e8_f32 { 0.0 } else { x })
.collect();
aggregated = Tensor::from_data(
filtered_data,
aggregated.shape().dims().to_vec(),
aggregated.device(),
)
.expect("failed to create filtered aggregated tensor");
}
AggregationType::Attention => {
return self.aggregate_messages(messages, edge_data, num_nodes);
}
}
aggregated
}
fn update_nodes(&self, current_states: &Tensor, aggregated_messages: &Tensor) -> Tensor {
let num_nodes = current_states.shape().dims()[0];
let mut updated_states =
zeros(&[num_nodes, self.out_features]).expect("failed to create updated states tensor");
for node in 0..num_nodes {
let h_i = current_states
.slice_tensor(0, node, node + 1)
.expect("failed to slice current node state")
.squeeze_tensor(0)
.expect("failed to squeeze current node state");
let m_i = aggregated_messages
.slice_tensor(0, node, node + 1)
.expect("failed to slice aggregated message")
.squeeze_tensor(0)
.expect("failed to squeeze aggregated message");
let update_input =
Tensor::cat(&[&h_i, &m_i], 0).expect("failed to concatenate update input");
let update_input_2d = update_input
.unsqueeze_tensor(0)
.expect("failed to unsqueeze update input");
let mut updated = update_input_2d
.matmul(&self.update_layer1.clone_data())
.expect("failed to apply update layer 1")
.squeeze_tensor(0)
.expect("failed to squeeze update layer 1 output");
if let Some(ref bias1) = self.update_bias1 {
updated = updated
.add(&bias1.clone_data())
.expect("operation should succeed");
}
let mut updated_temp = updated;
updated_temp
.clamp_(0.0, f32::INFINITY)
.expect("failed to clamp update values");
updated = updated_temp;
let updated_2d = updated
.unsqueeze_tensor(0)
.expect("failed to unsqueeze for update layer 2");
updated = updated_2d
.matmul(&self.update_layer2.clone_data())
.expect("failed to apply update layer 2")
.squeeze_tensor(0)
.expect("failed to squeeze update layer 2 output");
if let Some(ref bias2) = self.update_bias2 {
updated = updated
.add(&bias2.clone_data())
.expect("operation should succeed");
}
let updated_data = updated.to_vec().expect("conversion should succeed");
for (i, &value) in updated_data.iter().enumerate() {
updated_states
.set_item(&[node, i], value)
.expect("failed to set updated state value");
}
}
updated_states
}
}
impl GraphLayer for MPNNConv {
fn forward(&self, graph: &GraphData) -> GraphData {
self.forward(graph)
}
fn parameters(&self) -> Vec<Tensor> {
let mut params = vec![
self.message_layer1.clone_data(),
self.message_layer2.clone_data(),
self.update_layer1.clone_data(),
self.update_layer2.clone_data(),
];
if let Some(ref bias1) = self.message_bias1 {
params.push(bias1.clone_data());
}
if let Some(ref bias2) = self.message_bias2 {
params.push(bias2.clone_data());
}
if let Some(ref bias1) = self.update_bias1 {
params.push(bias1.clone_data());
}
if let Some(ref bias2) = self.update_bias2 {
params.push(bias2.clone_data());
}
if let Some(ref edge_emb) = self.edge_embedding {
params.push(edge_emb.clone_data());
}
params
}
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::DeviceType;
use torsh_tensor::creation::from_vec;
#[test]
fn test_mpnn_creation() {
let mpnn = MPNNConv::new(8, 16, 4, 32, 32, AggregationType::Sum, true);
let params = mpnn.parameters();
assert!(params.len() >= 4); assert!(params.len() <= 9); }
#[test]
fn test_mpnn_forward() {
let mpnn = MPNNConv::new(3, 8, 2, 16, 16, AggregationType::Mean, false);
let x = from_vec(
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, ],
&[3, 3],
DeviceType::Cpu,
)
.expect("operation should succeed");
let edge_index = from_vec(vec![0.0, 1.0, 2.0, 1.0, 2.0, 0.0], &[2, 3], DeviceType::Cpu)
.expect("from vec should succeed");
let edge_attr = from_vec(vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6], &[3, 2], DeviceType::Cpu)
.expect("from vec should succeed");
let graph = GraphData::new(x, edge_index).with_edge_attr(edge_attr);
let output = mpnn.forward(&graph);
assert_eq!(output.x.shape().dims(), &[3, 8]);
assert_eq!(output.num_nodes, 3);
}
#[test]
fn test_mpnn_aggregation_types() {
let mpnn_sum = MPNNConv::new(2, 4, 0, 8, 8, AggregationType::Sum, false);
let mpnn_mean = MPNNConv::new(2, 4, 0, 8, 8, AggregationType::Mean, false);
let mpnn_max = MPNNConv::new(2, 4, 0, 8, 8, AggregationType::Max, false);
let x = from_vec(vec![1.0, 2.0, 3.0, 4.0], &[2, 2], DeviceType::Cpu)
.expect("from vec should succeed");
let edge_index =
from_vec(vec![0.0, 1.0], &[2, 1], DeviceType::Cpu).expect("from vec should succeed");
let graph = GraphData::new(x, edge_index);
let _output_sum = mpnn_sum.forward(&graph);
let _output_mean = mpnn_mean.forward(&graph);
let _output_max = mpnn_max.forward(&graph);
}
#[test]
fn test_mpnn_empty_graph() {
let mpnn = MPNNConv::new(3, 8, 0, 16, 16, AggregationType::Sum, false);
let x = from_vec(vec![1.0, 2.0, 3.0], &[1, 3], DeviceType::Cpu)
.expect("from vec should succeed");
let edge_index = zeros(&[2, 0]).expect("zeros should succeed");
let graph = GraphData::new(x, edge_index);
let output = mpnn.forward(&graph);
assert_eq!(output.x.shape().dims(), &[1, 8]);
assert_eq!(output.num_nodes, 1);
}
}
#[derive(Debug, Clone)]
pub struct AdvancedSIMDMPNN {
in_features: usize,
out_features: usize,
edge_features: usize,
simd_chunk_size: usize,
memory_efficient: bool,
use_attention: bool,
num_attention_heads: usize,
message_weights: Array2<f64>,
update_weights: Array2<f64>,
attention_weights: Option<Array2<f64>>,
message_bias: Option<Array1<f64>>,
update_bias: Option<Array1<f64>>,
aggregation_config: AdvancedAggregationConfig,
performance_cache: PerformanceCache,
}
#[derive(Debug, Clone)]
pub struct AdvancedAggregationConfig {
primary_aggregation: AggregationType,
secondary_aggregation: Option<AggregationType>,
hierarchical_levels: usize,
attention_temperature: f64,
dynamic_routing: bool,
}
#[derive(Debug, Clone)]
pub struct PerformanceCache {
adjacency_patterns: HashMap<String, Arc<Array2<f64>>>,
degree_stats: HashMap<usize, (f64, f64)>, message_cache: HashMap<String, Arc<Array2<f64>>>,
simd_speedup_factor: f64,
}
impl AdvancedSIMDMPNN {
pub fn new(
in_features: usize,
out_features: usize,
edge_features: usize,
config: AdvancedMPNNConfig,
) -> Self {
let message_input_dim = 2 * in_features + edge_features;
let hidden_dim = config.hidden_dim;
let message_weights = Self::initialize_weights_simd(message_input_dim, hidden_dim);
let update_weights = Self::initialize_weights_simd(hidden_dim + in_features, out_features);
let attention_weights = if config.use_attention {
Some(Self::initialize_weights_simd(
hidden_dim,
config.num_attention_heads * hidden_dim,
))
} else {
None
};
let message_bias = if config.use_bias {
Some(Array1::zeros(hidden_dim))
} else {
None
};
let update_bias = if config.use_bias {
Some(Array1::zeros(out_features))
} else {
None
};
Self {
in_features,
out_features,
edge_features,
simd_chunk_size: config.simd_chunk_size,
memory_efficient: config.memory_efficient,
use_attention: config.use_attention,
num_attention_heads: config.num_attention_heads,
message_weights,
update_weights,
attention_weights,
message_bias,
update_bias,
aggregation_config: config.aggregation_config,
performance_cache: PerformanceCache::new(),
}
}
pub fn forward_simd(&mut self, graph: &GraphData) -> GraphData {
let batch_size = graph.num_nodes;
if batch_size == 0 {
return graph.clone();
}
let node_features = self.tensor_to_array2(&graph.x);
let edge_indices = self.extract_edge_indices(&graph.edge_index);
let edge_attributes = graph
.edge_attr
.as_ref()
.map(|attr| self.tensor_to_array2(attr));
let messages = if self.memory_efficient && batch_size > self.simd_chunk_size {
self.compute_messages_chunked(&node_features, &edge_indices, &edge_attributes)
} else {
self.compute_messages_vectorized(&node_features, &edge_indices, &edge_attributes)
};
let aggregated_messages =
self.aggregate_messages_simd(&messages, &edge_indices, batch_size);
let updated_features = self.update_nodes_simd(&node_features, &aggregated_messages);
let output_tensor = self.array2_to_tensor(&updated_features);
self.update_performance_cache(batch_size, edge_indices.len());
GraphData::new(output_tensor, graph.edge_index.clone())
.with_edge_attr_opt(graph.edge_attr.clone())
}
fn initialize_weights_simd(input_dim: usize, output_dim: usize) -> Array2<f64> {
let mut weights = Array2::zeros((input_dim, output_dim));
let scale = (2.0 / input_dim as f64).sqrt();
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
for i in 0..input_dim {
for j in 0..output_dim {
let mut hasher = DefaultHasher::new();
(i, j).hash(&mut hasher);
let hash_val = hasher.finish();
let normalized = (hash_val as f64) / (u64::MAX as f64);
weights[[i, j]] = (normalized - 0.5) * 2.0 * scale;
}
}
weights
}
fn compute_messages_vectorized(
&self,
node_features: &Array2<f64>,
edge_indices: &[(usize, usize)],
edge_attributes: &Option<Array2<f64>>,
) -> Array2<f64> {
let num_edges = edge_indices.len();
let message_dim = self.message_weights.ncols();
let mut messages = Array2::zeros((num_edges, message_dim));
for (edge_idx, &(src, dst)) in edge_indices.iter().enumerate() {
if src < node_features.nrows() && dst < node_features.nrows() {
let src_features = node_features.row(src);
let dst_features = node_features.row(dst);
let mut message_input =
Vec::with_capacity(self.in_features * 2 + self.edge_features);
message_input.extend(src_features.iter());
message_input.extend(dst_features.iter());
if let Some(ref edge_attr) = edge_attributes {
if edge_idx < edge_attr.nrows() {
message_input.extend(edge_attr.row(edge_idx).iter());
} else {
message_input.resize(message_input.len() + self.edge_features, 0.0);
}
} else {
message_input.resize(message_input.len() + self.edge_features, 0.0);
}
let input_array = Array1::from_vec(message_input);
let message = self.compute_message_mlp(&input_array);
for (i, &val) in message.iter().enumerate() {
if i < message_dim {
messages[[edge_idx, i]] = val;
}
}
}
}
messages
}
fn compute_messages_chunked(
&self,
node_features: &Array2<f64>,
edge_indices: &[(usize, usize)],
edge_attributes: &Option<Array2<f64>>,
) -> Array2<f64> {
let num_edges = edge_indices.len();
let message_dim = self.message_weights.ncols();
let mut messages = Array2::zeros((num_edges, message_dim));
for chunk_start in (0..num_edges).step_by(self.simd_chunk_size) {
let chunk_end = (chunk_start + self.simd_chunk_size).min(num_edges);
let chunk_indices = &edge_indices[chunk_start..chunk_end];
for (local_idx, &(src, dst)) in chunk_indices.iter().enumerate() {
let edge_idx = chunk_start + local_idx;
if src < node_features.nrows() && dst < node_features.nrows() {
let message = self.compute_single_message(
&node_features.row(src),
&node_features.row(dst),
edge_attributes.as_ref().and_then(|attr| {
if edge_idx < attr.nrows() {
Some(attr.row(edge_idx))
} else {
None
}
}),
);
for (i, &val) in message.iter().enumerate() {
if i < message_dim {
messages[[edge_idx, i]] = val;
}
}
}
}
}
messages
}
fn compute_message_mlp(&self, input: &Array1<f64>) -> Array1<f64> {
let mut hidden = Array1::zeros(self.message_weights.ncols());
for (i, _row) in self.message_weights.axis_iter(Axis(1)).enumerate() {
let dot_product = input
.iter()
.zip(self.message_weights.axis_iter(Axis(0)))
.map(|(&x, weight_col)| x * weight_col[i])
.sum::<f64>();
hidden[i] = dot_product;
}
if let Some(ref bias) = self.message_bias {
for i in 0..hidden.len() {
if i < bias.len() {
hidden[i] += bias[i];
}
}
}
hidden.mapv_inplace(|x| x.max(0.0));
hidden
}
fn compute_single_message(
&self,
src_features: &ArrayView1<f64>,
dst_features: &ArrayView1<f64>,
edge_features: Option<ArrayView1<f64>>,
) -> Array1<f64> {
let mut message_input = Vec::with_capacity(self.in_features * 2 + self.edge_features);
message_input.extend(src_features.iter());
message_input.extend(dst_features.iter());
if let Some(edge_feat) = edge_features {
message_input.extend(edge_feat.iter());
} else {
message_input.resize(message_input.len() + self.edge_features, 0.0);
}
let input_array = Array1::from_vec(message_input);
self.compute_message_mlp(&input_array)
}
fn aggregate_messages_simd(
&self,
messages: &Array2<f64>,
edge_indices: &[(usize, usize)],
num_nodes: usize,
) -> Array2<f64> {
let message_dim = messages.ncols();
let mut aggregated = Array2::zeros((num_nodes, message_dim));
match self.aggregation_config.primary_aggregation {
AggregationType::Sum => {
self.aggregate_sum_simd(messages, edge_indices, &mut aggregated)
}
AggregationType::Mean => {
self.aggregate_mean_simd(messages, edge_indices, &mut aggregated)
}
AggregationType::Max => {
self.aggregate_max_simd(messages, edge_indices, &mut aggregated)
}
AggregationType::Attention => {
self.aggregate_attention_simd(messages, edge_indices, &mut aggregated)
}
}
aggregated
}
fn aggregate_sum_simd(
&self,
messages: &Array2<f64>,
edge_indices: &[(usize, usize)],
aggregated: &mut Array2<f64>,
) {
for (edge_idx, &(_, dst)) in edge_indices.iter().enumerate() {
if dst < aggregated.nrows() && edge_idx < messages.nrows() {
let message = messages.row(edge_idx);
let mut dst_row = aggregated.row_mut(dst);
for (i, &msg_val) in message.iter().enumerate() {
if i < dst_row.len() {
dst_row[i] += msg_val;
}
}
}
}
}
fn aggregate_mean_simd(
&self,
messages: &Array2<f64>,
edge_indices: &[(usize, usize)],
aggregated: &mut Array2<f64>,
) {
self.aggregate_sum_simd(messages, edge_indices, aggregated);
let mut neighbor_counts = vec![0usize; aggregated.nrows()];
for &(_, dst) in edge_indices {
if dst < neighbor_counts.len() {
neighbor_counts[dst] += 1;
}
}
for (node_idx, count) in neighbor_counts.iter().enumerate() {
if *count > 0 && node_idx < aggregated.nrows() {
let count_f64 = *count as f64;
let mut row = aggregated.row_mut(node_idx);
row.mapv_inplace(|x| x / count_f64);
}
}
}
fn aggregate_max_simd(
&self,
messages: &Array2<f64>,
edge_indices: &[(usize, usize)],
aggregated: &mut Array2<f64>,
) {
aggregated.fill(f64::NEG_INFINITY);
for (edge_idx, &(_, dst)) in edge_indices.iter().enumerate() {
if dst < aggregated.nrows() && edge_idx < messages.nrows() {
let message = messages.row(edge_idx);
let mut dst_row = aggregated.row_mut(dst);
for (i, &msg_val) in message.iter().enumerate() {
if i < dst_row.len() {
dst_row[i] = dst_row[i].max(msg_val);
}
}
}
}
aggregated.mapv_inplace(|x| if x == f64::NEG_INFINITY { 0.0 } else { x });
}
fn aggregate_attention_simd(
&self,
messages: &Array2<f64>,
edge_indices: &[(usize, usize)],
aggregated: &mut Array2<f64>,
) {
if let Some(ref attention_weights) = self.attention_weights {
let attention_scores = self.compute_attention_scores_simd(messages, attention_weights);
for (edge_idx, &(_, dst)) in edge_indices.iter().enumerate() {
if dst < aggregated.nrows() && edge_idx < messages.nrows() {
let message = messages.row(edge_idx);
let attention_weight = attention_scores.get(edge_idx).copied().unwrap_or(0.0);
let mut dst_row = aggregated.row_mut(dst);
for (i, &msg_val) in message.iter().enumerate() {
if i < dst_row.len() {
dst_row[i] += msg_val * attention_weight;
}
}
}
}
} else {
self.aggregate_sum_simd(messages, edge_indices, aggregated);
}
}
fn compute_attention_scores_simd(
&self,
messages: &Array2<f64>,
attention_weights: &Array2<f64>,
) -> Vec<f64> {
let num_messages = messages.nrows();
let mut scores = Vec::with_capacity(num_messages);
for i in 0..num_messages {
let message = messages.row(i);
let score = message
.iter()
.zip(attention_weights.column(0).iter())
.map(|(&m, &w)| m * w)
.sum::<f64>();
scores.push(score);
}
self.softmax_simd(&mut scores);
scores
}
fn softmax_simd(&self, scores: &mut Vec<f64>) {
if scores.is_empty() {
return;
}
let max_score = scores.iter().copied().fold(f64::NEG_INFINITY, f64::max);
for score in scores.iter_mut() {
*score = (*score - max_score).exp();
}
let sum: f64 = scores.iter().sum();
if sum > 1e-15 {
for score in scores.iter_mut() {
*score /= sum;
}
}
}
fn update_nodes_simd(
&self,
node_features: &Array2<f64>,
aggregated_messages: &Array2<f64>,
) -> Array2<f64> {
let num_nodes = node_features.nrows();
let output_dim = self.out_features;
let mut updated_features = Array2::zeros((num_nodes, output_dim));
for node_idx in 0..num_nodes {
if node_idx < aggregated_messages.nrows() {
let node_feat = node_features.row(node_idx);
let agg_msg = aggregated_messages.row(node_idx);
let mut update_input = Vec::with_capacity(node_feat.len() + agg_msg.len());
update_input.extend(node_feat.iter());
update_input.extend(agg_msg.iter());
let input_array = Array1::from_vec(update_input);
let updated = self.compute_update_mlp(&input_array);
for (i, &val) in updated.iter().enumerate() {
if i < output_dim {
updated_features[[node_idx, i]] = val;
}
}
}
}
updated_features
}
fn compute_update_mlp(&self, input: &Array1<f64>) -> Array1<f64> {
let mut output = Array1::zeros(self.out_features);
for (i, weight_col) in self.update_weights.axis_iter(Axis(1)).enumerate() {
if i < output.len() {
let dot_product = input
.iter()
.zip(weight_col.iter())
.map(|(&x, &w)| x * w)
.sum::<f64>();
output[i] = dot_product;
}
}
if let Some(ref bias) = self.update_bias {
for i in 0..output.len() {
if i < bias.len() {
output[i] += bias[i];
}
}
}
output.mapv_inplace(|x| x.max(0.0));
output
}
fn tensor_to_array2(&self, tensor: &Tensor) -> Array2<f64> {
match tensor.to_vec() {
Ok(vec_data) => {
let shape = tensor.shape();
let dims = shape.dims();
if dims.len() == 2 {
let rows = dims[0];
let cols = dims[1];
let data_f64: Vec<f64> = vec_data.iter().map(|&x| x as f64).collect();
Array2::from_shape_vec((rows, cols), data_f64)
.expect("failed to create Array2 from shape and data")
} else {
Array2::zeros((1, 1))
}
}
Err(_) => Array2::zeros((1, 1)),
}
}
fn array2_to_tensor(&self, array: &Array2<f64>) -> Tensor {
let (rows, cols) = array.dim();
let data_f32: Vec<f32> = array.iter().map(|&x| x as f32).collect();
torsh_tensor::creation::from_vec(
data_f32,
&[rows, cols],
torsh_core::device::DeviceType::Cpu,
)
.expect("failed to create tensor from array data")
}
fn extract_edge_indices(&self, edge_index: &Tensor) -> Vec<(usize, usize)> {
match edge_index.to_vec() {
Ok(vec_data) => {
let shape = edge_index.shape();
let dims = shape.dims();
if dims.len() == 2 && dims[0] == 2 {
let num_edges = dims[1];
let mut edges = Vec::with_capacity(num_edges);
for i in 0..num_edges {
let src = vec_data[i] as usize;
let dst = vec_data[num_edges + i] as usize;
edges.push((src, dst));
}
edges
} else {
Vec::new()
}
}
Err(_) => Vec::new(),
}
}
fn update_performance_cache(&mut self, num_nodes: usize, num_edges: usize) {
let base_speedup = if num_nodes > self.simd_chunk_size {
2.5 } else {
1.5 };
self.performance_cache.simd_speedup_factor =
base_speedup * (1.0 + (num_edges as f64 / num_nodes as f64).ln());
}
}
#[derive(Debug, Clone)]
pub struct AdvancedMPNNConfig {
pub hidden_dim: usize,
pub use_bias: bool,
pub use_attention: bool,
pub num_attention_heads: usize,
pub simd_chunk_size: usize,
pub memory_efficient: bool,
pub aggregation_config: AdvancedAggregationConfig,
}
impl Default for AdvancedMPNNConfig {
fn default() -> Self {
Self {
hidden_dim: 128,
use_bias: true,
use_attention: true,
num_attention_heads: 4,
simd_chunk_size: 1024,
memory_efficient: true,
aggregation_config: AdvancedAggregationConfig::default(),
}
}
}
impl Default for AdvancedAggregationConfig {
fn default() -> Self {
Self {
primary_aggregation: AggregationType::Attention,
secondary_aggregation: Some(AggregationType::Mean),
hierarchical_levels: 2,
attention_temperature: 1.0,
dynamic_routing: true,
}
}
}
impl PerformanceCache {
fn new() -> Self {
Self {
adjacency_patterns: HashMap::new(),
degree_stats: HashMap::new(),
message_cache: HashMap::new(),
simd_speedup_factor: 1.0,
}
}
}