#![allow(dead_code)]
use crate::parameter::Parameter;
use crate::GraphData;
use torsh_tensor::{
creation::{from_vec, randn, zeros},
Tensor,
};
pub mod global {
use super::*;
pub fn global_mean_pool(graph: &GraphData) -> Tensor {
graph
.x
.mean(Some(&[0]), false)
.expect("mean pooling should succeed")
}
pub fn global_max_pool(graph: &GraphData) -> Tensor {
graph
.x
.max(Some(0), false)
.expect("max pooling should succeed")
}
pub fn global_sum_pool(graph: &GraphData) -> Tensor {
graph
.x
.sum_dim(&[0], false)
.expect("sum pooling should succeed")
}
pub struct GlobalAttentionPool {
gate_nn: Parameter,
feat_nn: Parameter,
}
impl GlobalAttentionPool {
pub fn new(input_dim: usize, hidden_dim: usize) -> Self {
let gate_nn = Parameter::new(
randn(&[input_dim, hidden_dim]).expect("randn gate_nn should succeed"),
);
let feat_nn = Parameter::new(
randn(&[input_dim, hidden_dim]).expect("randn feat_nn should succeed"),
);
Self { gate_nn, feat_nn }
}
pub fn forward(&self, graph: &GraphData) -> Tensor {
let gate = graph
.x
.matmul(&self.gate_nn.clone_data())
.expect("operation should succeed")
.sigmoid()
.expect("sigmoid should succeed");
let feat = graph
.x
.matmul(&self.feat_nn.clone_data())
.expect("operation should succeed");
let weighted_features = feat.mul(&gate).expect("operation should succeed");
weighted_features
.sum_dim(&[0], false)
.expect("sum reduction should succeed")
}
pub fn parameters(&self) -> Vec<Tensor> {
vec![self.gate_nn.clone_data(), self.feat_nn.clone_data()]
}
}
pub struct Set2Set {
input_dim: usize,
hidden_dim: usize,
num_layers: usize,
num_iters: usize,
lstm_weights: Vec<Parameter>,
attention_weights: Parameter,
projection_weights: Parameter,
}
impl Set2Set {
pub fn new(
input_dim: usize,
hidden_dim: usize,
num_layers: usize,
num_iters: usize,
) -> Self {
let mut lstm_weights = Vec::new();
for _ in 0..num_layers {
lstm_weights.push(Parameter::new(
randn(&[hidden_dim * 4, hidden_dim + input_dim])
.expect("randn lstm weights should succeed"),
));
}
let attention_weights = Parameter::new(
randn(&[hidden_dim, input_dim]).expect("randn attention weights should succeed"),
);
let projection_weights = Parameter::new(
randn(&[input_dim, hidden_dim]).expect("randn projection weights should succeed"),
);
Self {
input_dim,
hidden_dim,
num_layers,
num_iters,
lstm_weights,
attention_weights,
projection_weights,
}
}
pub fn forward(&self, graph: &GraphData) -> Tensor {
let _num_nodes = graph.num_nodes;
let mut query = zeros(&[1, self.hidden_dim]).expect("zeros query should succeed");
for _ in 0..self.num_iters {
let scores = query
.matmul(&self.attention_weights.clone_data())
.expect("operation should succeed")
.matmul(&graph.x.t().expect("transpose should succeed"))
.expect("operation should succeed")
.softmax(-1)
.expect("softmax should succeed");
let attended = scores.matmul(&graph.x).expect("operation should succeed");
let projected_attended = attended
.matmul(&self.projection_weights.clone_data())
.expect("operation should succeed");
query = query
.add(&projected_attended)
.expect("operation should succeed");
}
query.squeeze(0).expect("squeeze should succeed")
}
pub fn parameters(&self) -> Vec<Tensor> {
let mut params: Vec<Tensor> =
self.lstm_weights.iter().map(|p| p.clone_data()).collect();
params.push(self.attention_weights.clone_data());
params.push(self.projection_weights.clone_data());
params
}
}
}
pub mod hierarchical {
use super::*;
pub struct DiffPool {
embed_dim: usize,
assign_dim: usize,
embed_gnn: Parameter,
assign_gnn: Parameter,
link_pred_loss_weight: f64,
entropy_loss_weight: f64,
}
impl DiffPool {
pub fn new(embed_dim: usize, assign_dim: usize) -> Self {
let embed_gnn = Parameter::new(
randn(&[embed_dim, embed_dim]).expect("randn embed_gnn should succeed"),
);
let assign_gnn = Parameter::new(
randn(&[embed_dim, assign_dim]).expect("randn assign_gnn should succeed"),
);
Self {
embed_dim,
assign_dim,
embed_gnn,
assign_gnn,
link_pred_loss_weight: 1.0,
entropy_loss_weight: 1.0,
}
}
pub fn forward(&self, graph: &GraphData) -> (GraphData, Tensor) {
let num_nodes = graph.num_nodes;
let node_embeddings = graph
.x
.matmul(&self.embed_gnn.clone_data())
.expect("operation should succeed");
let assignment_logits = graph
.x
.matmul(&self.assign_gnn.clone_data())
.expect("operation should succeed");
let assignment_matrix = assignment_logits
.softmax(-1)
.expect("softmax should succeed");
let pooled_features = assignment_matrix
.t()
.expect("transpose should succeed")
.matmul(&node_embeddings)
.expect("operation should succeed");
let adjacency = self.compute_adjacency_matrix(&graph.edge_index, num_nodes);
let pooled_adj = assignment_matrix
.t()
.expect("transpose should succeed")
.matmul(&adjacency)
.expect("operation should succeed")
.matmul(&assignment_matrix)
.expect("operation should succeed");
let (new_edge_index, _) = self.adjacency_to_edge_index(&pooled_adj);
let link_pred_loss = self.compute_link_prediction_loss(&adjacency, &assignment_matrix);
let entropy_loss = self.compute_entropy_loss(&assignment_matrix);
let total_aux_loss = link_pred_loss
.mul_scalar(self.link_pred_loss_weight as f32)
.expect("mul_scalar link_pred should succeed")
.add(
&entropy_loss
.mul_scalar(self.entropy_loss_weight as f32)
.expect("mul_scalar entropy should succeed"),
)
.expect("operation should succeed");
let pooled_graph = GraphData {
x: pooled_features,
edge_index: new_edge_index,
edge_attr: None,
batch: None,
num_nodes: self.assign_dim,
num_edges: 0, };
(pooled_graph, total_aux_loss)
}
fn compute_adjacency_matrix(&self, edge_index: &Tensor, num_nodes: usize) -> Tensor {
let mut adjacency =
zeros(&[num_nodes, num_nodes]).expect("zeros adjacency should succeed");
let edge_data = edge_index.to_vec().expect("conversion should succeed");
let edge_list: Vec<Vec<i64>> = vec![
edge_data[0..edge_data.len() / 2]
.iter()
.map(|&x| x as i64)
.collect(),
edge_data[edge_data.len() / 2..]
.iter()
.map(|&x| x as i64)
.collect(),
];
for j in 0..edge_list[0].len() {
let src = edge_list[0][j] as usize;
let dst = edge_list[1][j] as usize;
if src < num_nodes && dst < num_nodes {
let mut adj_data = adjacency.to_vec().expect("conversion should succeed");
adj_data[src * num_nodes + dst] = 1.0;
adjacency = torsh_tensor::creation::from_vec(
adj_data,
&[num_nodes, num_nodes],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec adjacency should succeed");
}
}
adjacency
}
fn adjacency_to_edge_index(&self, adjacency: &Tensor) -> (Tensor, usize) {
let adj_data = adjacency.to_vec().expect("conversion should succeed");
let mut edges = Vec::new();
let shape = adjacency.shape();
let (rows, cols) = (shape.dims()[0], shape.dims()[1]);
for i in 0..rows {
for j in 0..cols {
let idx = i * cols + j;
if idx < adj_data.len() && adj_data[idx] > 0.5 {
edges.push([i as f32, j as f32]);
}
}
}
if edges.is_empty() {
(zeros(&[2, 0]).expect("zeros empty edges should succeed"), 0)
} else {
let num_edges = edges.len();
let mut edge_vec = Vec::with_capacity(2 * num_edges);
for edge in &edges {
edge_vec.push(edge[0]);
}
for edge in &edges {
edge_vec.push(edge[1]);
}
(
from_vec(
edge_vec.iter().map(|&x| x as f32).collect(),
&[2, num_edges],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec edge_index should succeed"),
num_edges,
)
}
}
fn compute_link_prediction_loss(&self, adjacency: &Tensor, assignment: &Tensor) -> Tensor {
let predicted_adj = assignment
.matmul(&assignment.t().expect("transpose should succeed"))
.expect("operation should succeed");
let eps = 1e-8;
let eps_tensor = torsh_tensor::creation::ones_like(adjacency)
.expect("ones_like should succeed")
.mul_scalar(eps as f32)
.expect("mul_scalar eps should succeed");
let one_tensor =
torsh_tensor::creation::ones_like(adjacency).expect("ones_like should succeed");
let pos_loss = adjacency
.mul(
&predicted_adj
.add(&eps_tensor)
.expect("operation should succeed")
.ln()
.expect("ln should succeed"),
)
.expect("operation should succeed");
let neg_loss = one_tensor
.sub(adjacency)
.expect("operation should succeed")
.mul(
&one_tensor
.sub(&predicted_adj)
.expect("operation should succeed")
.add(&eps_tensor)
.expect("operation should succeed")
.ln()
.expect("ln should succeed"),
)
.expect("operation should succeed");
pos_loss
.add(&neg_loss)
.expect("operation should succeed")
.mean(None, false)
.expect("reduction should succeed")
.neg()
.expect("operation should succeed")
}
fn compute_entropy_loss(&self, assignment: &Tensor) -> Tensor {
let eps = 1e-8;
let eps_tensor = torsh_tensor::creation::ones_like(assignment)
.expect("ones_like should succeed")
.mul_scalar(eps as f32)
.expect("mul_scalar eps should succeed");
let entropy = assignment
.mul(
&assignment
.add(&eps_tensor)
.expect("operation should succeed")
.ln()
.expect("operation should succeed"),
)
.expect("operation should succeed")
.sum()
.expect("reduction should succeed")
.mean(None, false)
.expect("reduction should succeed")
.neg()
.expect("operation should succeed");
entropy
}
pub fn parameters(&self) -> Vec<Tensor> {
vec![self.embed_gnn.clone_data(), self.assign_gnn.clone_data()]
}
}
pub struct TopKPool {
ratio: f32,
min_score: Option<f32>,
score_layer: Parameter,
}
impl TopKPool {
pub fn new(input_dim: usize, ratio: f32, min_score: Option<f32>) -> Self {
let score_layer =
Parameter::new(randn(&[input_dim, 1]).expect("randn score_layer should succeed"));
Self {
ratio,
min_score,
score_layer,
}
}
pub fn forward(&self, graph: &GraphData) -> GraphData {
let num_nodes = graph.num_nodes;
let k = (num_nodes as f32 * self.ratio).ceil() as usize;
let scores = graph
.x
.matmul(&self.score_layer.clone_data())
.expect("operation should succeed")
.squeeze(-1)
.expect("squeeze should succeed");
let (top_scores, top_indices) = self.topk(&scores, k);
let (selected_indices, _selected_scores) = if let Some(min_score) = self.min_score {
let valid_mask = top_scores
.gt_scalar(min_score)
.expect("gt_scalar should succeed");
let mask_data = valid_mask.to_vec().expect("conversion should succeed");
let mask_f32 = mask_data
.iter()
.map(|&x| if x { 1.0 } else { 0.0 })
.collect();
let mask_tensor = from_vec(
mask_f32,
valid_mask.shape().dims(),
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec mask should succeed");
let valid_indices = self.masked_select(&top_indices, &mask_tensor);
let valid_scores = self.masked_select(&top_scores, &mask_tensor);
(valid_indices, valid_scores)
} else {
(top_indices, top_scores)
};
let selected_features = self.index_select(&graph.x, &selected_indices, 0);
let (new_edge_index, new_num_edges) =
self.filter_edges(&graph.edge_index, &selected_indices);
GraphData {
x: selected_features,
edge_index: new_edge_index,
edge_attr: graph.edge_attr.clone(), batch: None, num_nodes: selected_indices.shape().dims()[0],
num_edges: new_num_edges,
}
}
fn topk(&self, tensor: &Tensor, k: usize) -> (Tensor, Tensor) {
let values = tensor.to_vec().expect("conversion should succeed");
let mut indexed_values: Vec<(f32, usize)> = values
.into_iter()
.enumerate()
.map(|(i, v)| (v, i))
.collect();
indexed_values
.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
indexed_values.truncate(k);
let top_values: Vec<f32> = indexed_values.iter().map(|(v, _)| *v).collect();
let top_indices: Vec<f32> = indexed_values.iter().map(|(_, i)| *i as f32).collect();
let values_tensor = from_vec(top_values, &[k], torsh_core::device::DeviceType::Cpu)
.expect("from_vec values should succeed");
let indices_tensor = from_vec(top_indices, &[k], torsh_core::device::DeviceType::Cpu)
.expect("from_vec indices should succeed");
(values_tensor, indices_tensor)
}
fn masked_select(&self, tensor: &Tensor, mask: &Tensor) -> Tensor {
let values = tensor.to_vec().expect("conversion should succeed");
let mask_values = mask.to_vec().expect("conversion should succeed");
let selected: Vec<f32> = values
.into_iter()
.zip(mask_values.into_iter())
.filter_map(|(v, m)| if m > 0.5 { Some(v) } else { None })
.collect();
let selected_len = selected.len();
from_vec(
selected,
&[selected_len],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec selected should succeed")
}
fn index_select(&self, tensor: &Tensor, indices: &Tensor, dim: i64) -> Tensor {
let idx_values = indices.to_vec().expect("conversion should succeed");
if dim == 0 {
let tensor_data = tensor.to_vec().expect("conversion should succeed");
let shape = tensor.shape();
let cols = shape.dims()[1];
let original_data: Vec<Vec<f32>> = tensor_data
.chunks(cols)
.map(|chunk| chunk.to_vec())
.collect();
let mut selected_rows = Vec::new();
for &idx in &idx_values {
let idx_usize = idx as usize;
if idx_usize < original_data.len() {
selected_rows.extend_from_slice(&original_data[idx_usize]);
}
}
let num_rows = idx_values.len();
let num_cols = if num_rows > 0 {
selected_rows.len() / num_rows
} else {
0
};
from_vec(
selected_rows,
&[num_rows, num_cols],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec selected_rows should succeed")
} else {
tensor.clone()
}
}
fn filter_edges(&self, edge_index: &Tensor, selected_nodes: &Tensor) -> (Tensor, usize) {
let edge_data = edge_index.to_vec().expect("conversion should succeed");
let edges = vec![
edge_data[0..edge_data.len() / 2]
.iter()
.map(|&x| x as i64)
.collect::<Vec<i64>>(),
edge_data[edge_data.len() / 2..]
.iter()
.map(|&x| x as i64)
.collect::<Vec<i64>>(),
];
let selected_indices = selected_nodes.to_vec().expect("conversion should succeed");
let mut node_mapping = std::collections::HashMap::new();
for (new_idx, &old_idx) in selected_indices.iter().enumerate() {
node_mapping.insert(old_idx as i64, new_idx as i64);
}
let mut filtered_edges = Vec::new();
for j in 0..edges[0].len() {
let src = edges[0][j];
let dst = edges[1][j];
if let (Some(&new_src), Some(&new_dst)) =
(node_mapping.get(&src), node_mapping.get(&dst))
{
filtered_edges.push([new_src, new_dst]);
}
}
if filtered_edges.is_empty() {
(zeros(&[2, 0]).expect("zeros empty edges should succeed"), 0)
} else {
let num_edges = filtered_edges.len();
let mut edge_vec = Vec::with_capacity(2 * num_edges);
for edge in &filtered_edges {
edge_vec.push(edge[0]);
}
for edge in &filtered_edges {
edge_vec.push(edge[1]);
}
(
from_vec(
edge_vec.iter().map(|&x| x as f32).collect(),
&[2, num_edges],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec filtered edges should succeed"),
num_edges,
)
}
}
pub fn parameters(&self) -> Vec<Tensor> {
vec![self.score_layer.clone_data()]
}
}
pub struct MinCutPool {
input_dim: usize,
output_dim: usize,
assignment_layer: Parameter,
}
impl MinCutPool {
pub fn new(input_dim: usize, output_dim: usize) -> Self {
let assignment_layer = Parameter::new(
randn(&[input_dim, output_dim]).expect("randn assignment_layer should succeed"),
);
Self {
input_dim,
output_dim,
assignment_layer,
}
}
pub fn forward(&self, graph: &GraphData) -> (GraphData, Tensor) {
let assignment_logits = graph
.x
.matmul(&self.assignment_layer.clone_data())
.expect("operation should succeed");
let assignment_matrix = assignment_logits
.softmax(-1)
.expect("softmax should succeed");
let pooled_features = assignment_matrix
.t()
.expect("transpose should succeed")
.matmul(&graph.x)
.expect("operation should succeed");
let adjacency = self.compute_adjacency_matrix(&graph.edge_index, graph.num_nodes);
let pooled_adj = assignment_matrix
.t()
.expect("transpose should succeed")
.matmul(&adjacency)
.expect("operation should succeed")
.matmul(&assignment_matrix)
.expect("operation should succeed");
let (new_edge_index, new_num_edges) = self.adjacency_to_edge_index(&pooled_adj);
let mincut_loss = self.compute_mincut_loss(&adjacency, &assignment_matrix);
let orthogonality_loss = self.compute_orthogonality_loss(&assignment_matrix);
let total_loss = mincut_loss
.add(&orthogonality_loss)
.expect("operation should succeed");
let pooled_graph = GraphData {
x: pooled_features,
edge_index: new_edge_index,
edge_attr: None,
batch: None,
num_nodes: self.output_dim,
num_edges: new_num_edges,
};
(pooled_graph, total_loss)
}
fn compute_adjacency_matrix(&self, edge_index: &Tensor, num_nodes: usize) -> Tensor {
let mut adjacency =
zeros(&[num_nodes, num_nodes]).expect("zeros adjacency should succeed");
let edge_data = edge_index.to_vec().expect("conversion should succeed");
let edge_list: Vec<Vec<i64>> = vec![
edge_data[0..edge_data.len() / 2]
.iter()
.map(|&x| x as i64)
.collect(),
edge_data[edge_data.len() / 2..]
.iter()
.map(|&x| x as i64)
.collect(),
];
for j in 0..edge_list[0].len() {
let src = edge_list[0][j] as usize;
let dst = edge_list[1][j] as usize;
if src < num_nodes && dst < num_nodes {
let mut adj_data = adjacency.to_vec().expect("conversion should succeed");
adj_data[src * num_nodes + dst] = 1.0;
adjacency = torsh_tensor::creation::from_vec(
adj_data,
&[num_nodes, num_nodes],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec adjacency should succeed");
}
}
adjacency
}
fn adjacency_to_edge_index(&self, adjacency: &Tensor) -> (Tensor, usize) {
let adj_data = adjacency.to_vec().expect("conversion should succeed");
let mut edges = Vec::new();
let shape = adjacency.shape();
let (rows, cols) = (shape.dims()[0], shape.dims()[1]);
for i in 0..rows {
for j in 0..cols {
let idx = i * cols + j;
if idx < adj_data.len() && adj_data[idx] > 0.1 {
edges.push([i as f32, j as f32]);
}
}
}
if edges.is_empty() {
(zeros(&[2, 0]).expect("zeros empty edges should succeed"), 0)
} else {
let num_edges = edges.len();
let mut edge_vec = Vec::with_capacity(2 * num_edges);
for edge in &edges {
edge_vec.push(edge[0]);
}
for edge in &edges {
edge_vec.push(edge[1]);
}
(
from_vec(
edge_vec.iter().map(|&x| x as f32).collect(),
&[2, num_edges],
torsh_core::device::DeviceType::Cpu,
)
.expect("from_vec edge_index should succeed"),
num_edges,
)
}
}
fn compute_mincut_loss(&self, adjacency: &Tensor, assignment: &Tensor) -> Tensor {
let cut = assignment
.t()
.expect("transpose should succeed")
.matmul(adjacency)
.expect("operation should succeed")
.matmul(assignment)
.expect("operation should succeed");
let degree = assignment
.sum_dim(&[0], false)
.expect("sum_dim should succeed");
let degree_unsqueezed = degree.unsqueeze(0).expect("unsqueeze should succeed");
let degree_t = degree.unsqueeze(1).expect("unsqueeze should succeed");
let degree_product = degree_t
.matmul(°ree_unsqueezed)
.expect("operation should succeed");
let eps_tensor = torsh_tensor::creation::ones_like(°ree_product)
.expect("ones_like should succeed")
.mul_scalar(1e-8_f32)
.expect("mul_scalar eps should succeed");
let normalized_cut = cut
.div(
°ree_product
.add(&eps_tensor)
.expect("operation should succeed"),
)
.expect("operation should succeed");
let diag_sum = normalized_cut.sum().expect("reduction should succeed");
diag_sum.neg().expect("neg should succeed")
}
fn compute_orthogonality_loss(&self, assignment: &Tensor) -> Tensor {
let cluster_sizes = assignment.sum().expect("reduction should succeed");
let normalized_sizes = cluster_sizes
.div(&cluster_sizes.sum().expect("reduction should succeed"))
.expect("operation should succeed");
let eps = 1e-8;
let eps_tensor = torsh_tensor::creation::ones_like(&normalized_sizes)
.expect("ones_like should succeed")
.mul_scalar(eps as f32)
.expect("mul_scalar eps should succeed");
let entropy_loss = normalized_sizes
.mul(
&normalized_sizes
.add(&eps_tensor)
.expect("operation should succeed")
.ln()
.expect("ln should succeed"),
)
.expect("operation should succeed")
.sum()
.expect("reduction should succeed")
.neg()
.expect("neg should succeed");
entropy_loss.neg().expect("neg should succeed")
}
pub fn parameters(&self) -> Vec<Tensor> {
vec![self.assignment_layer.clone_data()]
}
}
}