use burn::nn::{Linear, LinearConfig};
use burn::tensor::backend::Backend;
use burn::tensor::Tensor;
use crate::hyperbolic::PoincareBall;
pub struct GCNConv<B: Backend> {
linear: Linear<B>,
}
impl<B: Backend> GCNConv<B> {
pub fn new(linear: Linear<B>) -> Self {
Self { linear }
}
pub fn init(d_in: usize, d_out: usize, device: &B::Device) -> Self {
Self {
linear: LinearConfig::new(d_in, d_out).init(device),
}
}
pub fn linear(&self) -> &Linear<B> {
&self.linear
}
pub fn forward(&self, x: Tensor<B, 2>, adj: Tensor<B, 2>) -> Tensor<B, 2> {
let x = self.linear.forward(x);
adj.matmul(x)
}
}
pub struct HGCNConv<B: Backend> {
linear: Linear<B>,
ball: PoincareBall,
}
impl<B: Backend> HGCNConv<B> {
pub fn new(linear: Linear<B>, c: f64) -> Self {
Self {
linear,
ball: PoincareBall::new(c),
}
}
pub fn init(d: usize, c: f64, device: &B::Device) -> Self {
Self {
linear: LinearConfig::new(d, d).init(device),
ball: PoincareBall::new(c),
}
}
pub fn linear(&self) -> &Linear<B> {
&self.linear
}
pub fn ball(&self) -> &PoincareBall {
&self.ball
}
pub fn forward(&self, x: Tensor<B, 2>, adj: Tensor<B, 2>) -> Tensor<B, 2> {
let x_tangent = self.ball.log0(x);
let x_tangent = self.linear.forward(x_tangent);
let aggregated = adj.matmul(x_tangent);
self.ball.exp0(aggregated)
}
pub fn forward_act<F>(
&self,
x: Tensor<B, 2>,
adj: Tensor<B, 2>,
act: F,
ball_out: &PoincareBall,
) -> Tensor<B, 2>
where
F: Fn(Tensor<B, 2>) -> Tensor<B, 2>,
{
let h = self.forward(x, adj);
self.ball.hyp_act(h, act, ball_out)
}
pub fn forward_with_basepoint(
&self,
x: Tensor<B, 2>,
adj: Tensor<B, 2>,
p: Tensor<B, 2>,
) -> Tensor<B, 2> {
let [n, d] = x.dims();
let [pn, _pd] = p.dims();
let p = if pn == 1 { p.expand([n, d]) } else { p };
let x_tangent = self.ball.log_map(p.clone(), x);
let x_tangent = self.linear.forward(x_tangent);
let aggregated = adj.matmul(x_tangent);
self.ball.exp_map(p, aggregated)
}
pub fn forward_with_basepoint_and_bias(
&self,
x: Tensor<B, 2>,
adj: Tensor<B, 2>,
p: Tensor<B, 2>,
b0: Tensor<B, 2>,
) -> Tensor<B, 2> {
let [n, d] = x.dims();
let [pn, _] = p.dims();
let [bn, _] = b0.dims();
let p = if pn == 1 { p.expand([n, d]) } else { p };
let b0 = if bn == 1 { b0.expand([n, d]) } else { b0 };
let x_tangent = self.ball.log_map(p.clone(), x);
let x_tangent = self.linear.forward(x_tangent);
let aggregated = adj.matmul(x_tangent);
let bias_p = self.ball.parallel_transport_0_to_x(p.clone(), b0);
let aggregated = aggregated + bias_p;
self.ball.exp_map(p, aggregated)
}
pub fn forward_local_dense(&self, x: Tensor<B, 2>, adj: Tensor<B, 2>) -> Tensor<B, 2> {
let [n, d] = x.dims();
let x = self.ball.project(x);
let p = x
.clone()
.reshape([n, 1, d])
.expand([n, n, d])
.reshape([n * n, d]);
let y = x
.clone()
.reshape([1, n, d])
.expand([n, n, d])
.reshape([n * n, d]);
let v = self.ball.log_map(p, y);
let v = self.linear.forward(v);
let v = v.reshape([n, n, d]);
let w = adj.reshape([n, n, 1]).expand([n, n, d]);
let agg = (v * w).sum_dim(1).reshape([n, d]);
self.ball.exp_map(x, agg)
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::TensorData;
use burn_ndarray::NdArray;
type B = NdArray<f32>;
fn dev() -> <B as Backend>::Device {
<B as Backend>::Device::default()
}
#[test]
fn gcn_forward_shapes() {
let n = 5;
let d = 3;
let layer = GCNConv::<B>::init(d, d, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.1f32; n * d], [n, d]), &dev());
let adj = Tensor::from_data(TensorData::new(vec![1.0f32; n * n], [n, n]), &dev());
let y = layer.forward(x, adj);
assert_eq!(y.dims(), [n, d]);
}
#[test]
fn hgcn_forward_shapes() {
let n = 5;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.01f32; n * d], [n, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y = layer.forward(x, adj);
assert_eq!(y.dims(), [n, d]);
}
#[test]
fn hgcn_with_basepoint_shapes() {
let n = 6;
let d = 4;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.01f32; n * d], [n, d]), &dev());
let p = Tensor::from_data(TensorData::new(vec![0.0f32; d], [1, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y = layer.forward_with_basepoint(x, adj, p);
assert_eq!(y.dims(), [n, d]);
}
#[test]
fn hgcn_local_dense_identity_adj_shapes() {
let n = 4;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::<B, 2>::from_data(TensorData::new(vec![0.01f32; n * d], [n, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y = layer.forward_local_dense(x, adj);
assert_eq!(y.dims(), [n, d]);
}
#[test]
fn hgcn_forward_act_produces_finite() {
let n = 4;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.05f32; n * d], [n, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let ball = *layer.ball();
let y = layer.forward_act(x, adj, |t| t.clamp_min(0.0), &ball);
assert_eq!(y.dims(), [n, d]);
let y_v = y.to_data().to_vec::<f32>().unwrap();
assert!(y_v.iter().all(|v| v.is_finite()), "forward_act non-finite");
}
#[test]
fn hgcn_forward_act_with_curvature_change() {
let n = 3;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(
TensorData::new(
vec![0.05f32, -0.03, 0.02, 0.01, 0.04, -0.01, -0.02, 0.01, 0.03],
[n, d],
),
&dev(),
);
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let ball_out = crate::PoincareBall::new(2.0);
let y = layer.forward_act(x, adj, |t| t.clamp_min(0.0), &ball_out);
assert_eq!(y.dims(), [n, d]);
let y_v = y.to_data().to_vec::<f32>().unwrap();
assert!(
y_v.iter().all(|v| v.is_finite()),
"curvature-change non-finite"
);
}
#[test]
fn hgcn_forward_with_bias_shapes_and_finite() {
let n = 4;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.05f32; n * d], [n, d]), &dev());
let b0 = Tensor::from_data(TensorData::new(vec![0.01f32, -0.01, 0.005], [1, d]), &dev());
let p = Tensor::from_data(TensorData::new(vec![0.0f32; d], [1, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y = layer.forward_with_basepoint_and_bias(x, adj, p, b0);
assert_eq!(y.dims(), [n, d]);
let y_v = y.to_data().to_vec::<f32>().unwrap();
assert!(
y_v.iter().all(|v| v.is_finite()),
"forward_with_bias non-finite"
);
}
#[test]
fn hgcn_forward_with_bias_differs_from_without() {
let n = 3;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(
TensorData::new(
vec![0.05f32, -0.03, 0.02, 0.01, 0.04, -0.01, -0.02, 0.01, 0.03],
[n, d],
),
&dev(),
);
let p = Tensor::from_data(TensorData::new(vec![0.0f32; d], [1, d]), &dev());
let b0 = Tensor::from_data(TensorData::new(vec![0.1f32, -0.05, 0.02], [1, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y_no_bias = layer.forward_with_basepoint(x.clone(), adj.clone(), p.clone());
let y_with_bias = layer.forward_with_basepoint_and_bias(x, adj, p, b0);
let a = y_no_bias.to_data().to_vec::<f32>().unwrap();
let b = y_with_bias.to_data().to_vec::<f32>().unwrap();
let diff: f32 = a.iter().zip(&b).map(|(x, y)| (x - y).abs()).sum();
assert!(diff > 1e-3, "bias should change output, diff={diff}");
}
#[test]
fn gcn_non_square_dimensions() {
let n = 4;
let d_in = 5;
let d_out = 3;
let layer = GCNConv::<B>::init(d_in, d_out, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.1f32; n * d_in], [n, d_in]), &dev());
let adj = Tensor::from_data(TensorData::new(vec![1.0f32; n * n], [n, n]), &dev());
let y = layer.forward(x, adj);
assert_eq!(y.dims(), [n, d_out]);
}
#[test]
fn hgcn_with_real_adjacency() {
let n = 4;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(
TensorData::new(
vec![
0.05f32, -0.03, 0.02, 0.01, 0.04, -0.01, -0.02, 0.01, 0.03, 0.03, -0.02, 0.01, ],
[n, d],
),
&dev(),
);
#[rustfmt::skip]
let adj_v = vec![
0.5, 0.5, 0.0, 0.0,
0.33, 0.33, 0.33, 0.0,
0.0, 0.33, 0.33, 0.33,
0.0, 0.0, 0.5, 0.5,
];
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y = layer.forward(x, adj);
assert_eq!(y.dims(), [n, d]);
let y_v = y.to_data().to_vec::<f32>().unwrap();
assert!(
y_v.iter().all(|v| v.is_finite()),
"chain graph forward non-finite"
);
let row0: Vec<f32> = y_v[0..d].to_vec();
let row1: Vec<f32> = y_v[d..2 * d].to_vec();
let diff: f32 = row0.iter().zip(&row1).map(|(a, b)| (a - b).abs()).sum();
assert!(
diff > 1e-4,
"boundary and interior nodes should differ, diff={diff}"
);
}
#[test]
fn two_layer_hgcn_pipeline() {
let n = 4;
let d = 3;
let layer1 = HGCNConv::<B>::init(d, 1.0, &dev());
let layer2 = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.05f32; n * d], [n, d]), &dev());
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
if i + 1 < n {
adj_v[i * n + i + 1] = 0.5;
adj_v[(i + 1) * n + i] = 0.5;
}
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let ball = *layer1.ball();
let h = layer1.forward_act(x, adj.clone(), |t| t.clamp_min(0.0), &ball);
let y = layer2.forward(h, adj);
assert_eq!(y.dims(), [n, d]);
let y_v = y.to_data().to_vec::<f32>().unwrap();
assert!(
y_v.iter().all(|v| v.is_finite()),
"two-layer pipeline non-finite"
);
}
#[test]
fn forward_local_dense_and_forward_agree_on_identity_adj() {
let n = 3;
let d = 3;
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::<B, 2>::from_data(
TensorData::new(
vec![0.05f32, -0.03, 0.02, 0.01, 0.04, -0.01, -0.02, 0.01, 0.03],
[n, d],
),
&dev(),
);
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y_origin = layer.forward(x.clone(), adj.clone());
let y_local = layer.forward_local_dense(x, adj);
let y_o = y_origin.to_data().to_vec::<f32>().unwrap();
let y_l = y_local.to_data().to_vec::<f32>().unwrap();
fn l1(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| (x - y).abs()).sum()
}
assert!(
y_o.iter().all(|v| v.is_finite()),
"forward produced non-finite"
);
assert!(
y_l.iter().all(|v| v.is_finite()),
"forward_local_dense produced non-finite"
);
assert!(
l1(&y_o, &y_l) < 0.5,
"forward vs forward_local_dense diverged: l1={}",
l1(&y_o, &y_l)
);
}
fn assert_inside_ball(t: &Tensor<B, 2>, ball: &crate::PoincareBall, label: &str) {
let [n, d] = t.dims();
let v = t.to_data().to_vec::<f32>().unwrap();
let max = ball.max_norm();
for i in 0..n {
let row = &v[i * d..(i + 1) * d];
let norm: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
norm <= max + 1e-4,
"{label} row {i} outside ball: norm={norm} max={max}"
);
assert!(
row.iter().all(|x| x.is_finite()),
"{label} row {i} non-finite"
);
}
}
#[test]
fn all_forward_variants_stay_inside_ball() {
let n = 4;
let d = 3;
let ball = crate::PoincareBall::new(1.0);
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(
TensorData::new(
vec![
0.3f32, -0.2, 0.1, 0.1, 0.2, -0.3, -0.1, 0.3, 0.2, 0.2, -0.1, 0.1,
],
[n, d],
),
&dev(),
);
#[rustfmt::skip]
let adj_v = vec![
0.5, 0.5, 0.0, 0.0,
0.33, 0.34, 0.33, 0.0,
0.0, 0.33, 0.34, 0.33,
0.0, 0.0, 0.5, 0.5f32,
];
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let p = Tensor::from_data(TensorData::new(vec![0.05f32, -0.02, 0.01], [1, d]), &dev());
let b0 = Tensor::from_data(TensorData::new(vec![0.02f32, -0.01, 0.005], [1, d]), &dev());
let y1 = layer.forward(x.clone(), adj.clone());
assert_inside_ball(&y1, &ball, "forward");
let y2 = layer.forward_with_basepoint(x.clone(), adj.clone(), p.clone());
assert_inside_ball(&y2, &ball, "forward_with_basepoint");
let y3 =
layer.forward_with_basepoint_and_bias(x.clone(), adj.clone(), p.clone(), b0.clone());
assert_inside_ball(&y3, &ball, "forward_with_basepoint_and_bias");
let y4 = layer.forward_local_dense(x.clone(), adj.clone());
assert_inside_ball(&y4, &ball, "forward_local_dense");
let y5 = layer.forward_act(x, adj, |t| t.clamp_min(0.0), &ball);
assert_inside_ball(&y5, &ball, "forward_act");
}
#[test]
fn per_node_basepoint() {
let n = 3;
let d = 3;
let ball = crate::PoincareBall::new(1.0);
let layer = HGCNConv::<B>::init(d, 1.0, &dev());
let x = Tensor::from_data(
TensorData::new(
vec![0.05f32, -0.03, 0.02, 0.01, 0.04, -0.01, -0.02, 0.01, 0.03],
[n, d],
),
&dev(),
);
let p = Tensor::from_data(
TensorData::new(
vec![0.02f32, 0.01, -0.01, -0.01, 0.02, 0.01, 0.01, -0.02, 0.02],
[n, d],
),
&dev(),
);
let mut adj_v = vec![0.0f32; n * n];
for i in 0..n {
adj_v[i * n + i] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(adj_v, [n, n]), &dev());
let y = layer.forward_with_basepoint(x, adj, p);
assert_eq!(y.dims(), [n, d]);
assert_inside_ball(&y, &ball, "per_node_basepoint");
}
}