use burn::module::{Ignored, Module, Param, ParamId};
use burn::nn::{Linear, LinearConfig};
use burn::tensor::backend::Backend;
use burn::tensor::{Distribution, IndexingUpdateOp, Int, Tensor};
use crate::hyperbolic::PoincareBall;
#[derive(Module, Debug)]
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)
}
}
#[derive(Module, Debug)]
pub struct HGCNConv<B: Backend> {
linear: Linear<B>,
ball: Ignored<PoincareBall>,
}
impl<B: Backend> HGCNConv<B> {
pub fn new(linear: Linear<B>, c: f64) -> Self {
Self {
linear,
ball: Ignored(PoincareBall::new(c)),
}
}
pub fn init(d: usize, c: f64, device: &B::Device) -> Self {
Self {
linear: LinearConfig::new(d, d).init(device),
ball: Ignored(PoincareBall::new(c)),
}
}
pub fn linear(&self) -> &Linear<B> {
&self.linear
}
pub fn ball(&self) -> &PoincareBall {
&self.ball.0
}
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)
}
}
#[derive(Module, Debug)]
pub struct RGCNConv<B: Backend> {
rel: Vec<Linear<B>>,
basis: Option<Param<Tensor<B, 3>>>,
coef: Option<Param<Tensor<B, 2>>>,
self_loop: Linear<B>,
}
impl<B: Backend> RGCNConv<B> {
pub fn init(d_in: usize, d_out: usize, num_relations: usize, device: &B::Device) -> Self {
Self {
rel: (0..num_relations)
.map(|_| LinearConfig::new(d_in, d_out).init(device))
.collect(),
basis: None,
coef: None,
self_loop: LinearConfig::new(d_in, d_out).init(device),
}
}
pub fn with_bases(
d_in: usize,
d_out: usize,
num_relations: usize,
num_bases: usize,
device: &B::Device,
) -> Self {
let std = (1.0 / d_in as f64).sqrt();
let mk3 = Tensor::random(
[num_bases, d_in, d_out],
Distribution::Normal(0.0, std),
device,
)
.require_grad();
let mk2 = Tensor::random(
[num_relations, num_bases],
Distribution::Normal(0.0, (1.0 / num_bases as f64).sqrt()),
device,
)
.require_grad();
Self {
rel: Vec::new(),
basis: Some(Param::initialized(ParamId::new(), mk3)),
coef: Some(Param::initialized(ParamId::new(), mk2)),
self_loop: LinearConfig::new(d_in, d_out).init(device),
}
}
pub fn num_relations(&self) -> usize {
match &self.coef {
Some(c) => c.val().dims()[0],
None => self.rel.len(),
}
}
pub fn forward(&self, x: Tensor<B, 2>, adjs: &[Tensor<B, 2>]) -> Tensor<B, 2> {
assert_eq!(
adjs.len(),
self.num_relations(),
"one adjacency per relation"
);
let mut out = self.self_loop.forward(x.clone());
if let (Some(basis), Some(coef)) = (&self.basis, &self.coef) {
let [nb, d_in, d_out] = basis.val().dims();
let flat = basis.val().reshape([nb, d_in * d_out]);
let ws = coef.val().matmul(flat); for (r, adj) in adjs.iter().enumerate() {
let w = ws
.clone()
.slice([r..r + 1, 0..d_in * d_out])
.reshape([d_in, d_out]);
out = out + adj.clone().matmul(x.clone().matmul(w));
}
} else {
for (lin, adj) in self.rel.iter().zip(adjs) {
out = out + adj.clone().matmul(lin.forward(x.clone()));
}
}
out
}
}
#[derive(Module, Debug)]
pub struct NBFConv<B: Backend> {
update: Linear<B>,
}
impl<B: Backend> NBFConv<B> {
pub fn init(d: usize, device: &B::Device) -> Self {
Self {
update: LinearConfig::new(d, d).init(device),
}
}
pub fn coverage(h: &Tensor<B, 2>) -> f32 {
let [n, d] = h.dims();
let v: Vec<f32> = h.clone().into_data().to_vec().unwrap();
let reached = (0..n)
.filter(|i| (0..d).any(|k| v[i * d + k] != 0.0))
.count();
reached as f32 / n.max(1) as f32
}
pub fn forward_edges(
&self,
h: Tensor<B, 3>,
h0: Tensor<B, 3>,
heads: Tensor<B, 1, Int>,
tails: Tensor<B, 1, Int>,
etypes: Tensor<B, 1, Int>,
rel: Tensor<B, 3>,
) -> Tensor<B, 3> {
let [q, n, d] = h.dims();
let src = h.select(1, heads); let w = rel.select(1, etypes); let msgs = src * w;
let agg = Tensor::zeros([q, n, d], &msgs.device()).select_assign(
1,
tails,
msgs,
IndexingUpdateOp::Add,
);
self.update.forward(agg + h0)
}
pub fn forward(
&self,
h: Tensor<B, 2>,
h0: Tensor<B, 2>,
adjs: &[Tensor<B, 2>],
rel: Tensor<B, 2>,
) -> Tensor<B, 2> {
let [num_types, d] = rel.dims();
assert_eq!(adjs.len(), num_types, "one adjacency per edge type");
let mut agg = h0;
for (t, adj) in adjs.iter().enumerate() {
let w = rel.clone().slice([t..t + 1, 0..d]); let scaled = h.clone() * w; agg = agg + adj.clone().matmul(scaled);
}
self.update.forward(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 rgcn_distinguishes_relations() {
let (n, d) = (4, 3);
let layer = RGCNConv::<B>::init(d, d, 2, &dev());
let x = Tensor::from_data(
TensorData::new((0..n * d).map(|i| i as f32 / 5.0).collect(), [n, d]),
&dev(),
);
let mut a_v = vec![0.0f32; n * n];
a_v[1] = 1.0; let mut b_v = vec![0.0f32; n * n];
b_v[2] = 1.0; let a = Tensor::from_data(TensorData::new(a_v, [n, n]), &dev());
let b = Tensor::from_data(TensorData::new(b_v, [n, n]), &dev());
let fwd: Vec<f32> = layer
.forward(x.clone(), &[a.clone(), b.clone()])
.into_data()
.to_vec()
.unwrap();
let swp: Vec<f32> = layer.forward(x, &[b, a]).into_data().to_vec().unwrap();
let diff: f32 = fwd.iter().zip(&swp).map(|(p, q)| (p - q).abs()).sum();
assert!(diff > 1e-4, "relation swap must change the output: {diff}");
}
#[test]
fn rgcn_zero_relations_is_self_loop() {
let (n, d) = (3, 2);
let layer = RGCNConv::<B>::init(d, d, 0, &dev());
let x = Tensor::from_data(TensorData::new(vec![0.3f32; n * d], [n, d]), &dev());
let y: Vec<f32> = layer.forward(x.clone(), &[]).into_data().to_vec().unwrap();
let s: Vec<f32> = layer.self_loop.forward(x).into_data().to_vec().unwrap();
assert_eq!(y, s);
}
#[test]
fn rgcn_basis_decomposition() {
let (n, d, r, nb) = (4, 3, 6, 2);
let layer = RGCNConv::<B>::with_bases(d, d, r, nb, &dev());
assert_eq!(layer.num_relations(), r);
let x = Tensor::from_data(TensorData::new(vec![0.2f32; n * d], [n, d]), &dev());
let eye = {
let mut v = vec![0.0f32; n * n];
for i in 0..n {
v[i * n + i] = 1.0;
}
Tensor::from_data(TensorData::new(v, [n, n]), &dev())
};
let adjs: Vec<_> = (0..r).map(|_| eye.clone()).collect();
let y = layer.forward(x, &adjs);
assert_eq!(y.dims(), [n, d]);
assert!(nb * d * d + r * nb < r * d * d);
}
#[test]
fn nbf_output_is_conditioned_on_the_source() {
let (n, d) = (4, 3);
let layer = NBFConv::<B>::init(d, &dev());
let mut ring = vec![0.0f32; n * n];
for i in 0..n {
ring[i * n + (i + 1) % n] = 1.0;
}
let adj = Tensor::from_data(TensorData::new(ring, [n, n]), &dev());
let rel = Tensor::from_data(TensorData::new(vec![0.5f32; d], [1, d]), &dev());
let indicator = |src: usize| {
let mut v = vec![0.0f32; n * d];
for j in 0..d {
v[src * d + j] = 1.0;
}
Tensor::<B, 2>::from_data(TensorData::new(v, [n, d]), &dev())
};
let h = Tensor::zeros([n, d], &dev());
let a: Vec<f32> = layer
.forward(
h.clone(),
indicator(0),
std::slice::from_ref(&adj),
rel.clone(),
)
.into_data()
.to_vec()
.unwrap();
let b: Vec<f32> = layer
.forward(h, indicator(2), &[adj], rel)
.into_data()
.to_vec()
.unwrap();
let diff: f32 = a.iter().zip(&b).map(|(p, q)| (p - q).abs()).sum();
assert!(diff > 1e-4, "indicator position must matter: {diff}");
}
#[test]
fn nbf_is_permutation_equivariant() {
let (n, d) = (3, 2);
let layer = NBFConv::<B>::init(d, &dev());
let adj_v = vec![0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0f32];
let adj = Tensor::from_data(TensorData::new(adj_v.clone(), [n, n]), &dev());
let h_v: Vec<f32> = (0..n * d).map(|i| i as f32 / 7.0).collect();
let h = Tensor::from_data(TensorData::new(h_v.clone(), [n, d]), &dev());
let h0 = Tensor::zeros([n, d], &dev());
let rel = Tensor::from_data(TensorData::new(vec![1.0f32; d], [1, d]), &dev());
let out: Vec<f32> = layer
.forward(h, h0.clone(), &[adj], rel.clone())
.into_data()
.to_vec()
.unwrap();
let sigma = [1usize, 2, 0];
let mut padj = vec![0.0f32; n * n];
for i in 0..n {
for j in 0..n {
if adj_v[i * n + j] == 1.0 {
padj[sigma[i] * n + sigma[j]] = 1.0;
}
}
}
let mut ph = vec![0.0f32; n * d];
for i in 0..n {
for k in 0..d {
ph[sigma[i] * d + k] = h_v[i * d + k];
}
}
let pout: Vec<f32> = layer
.forward(
Tensor::from_data(TensorData::new(ph, [n, d]), &dev()),
h0,
&[Tensor::from_data(TensorData::new(padj, [n, n]), &dev())],
rel,
)
.into_data()
.to_vec()
.unwrap();
for i in 0..n {
for k in 0..d {
let want = out[i * d + k];
let got = pout[sigma[i] * d + k];
assert!(
(want - got).abs() < 1e-5,
"output must permute with nodes at ({i},{k})"
);
}
}
}
#[test]
fn nbf_coverage_tracks_propagation() {
let (n, d) = (4, 2);
let mut path = vec![0.0f32; n * n];
for i in 0..n - 1 {
path[(i + 1) * n + i] = 1.0; }
let adj = Tensor::from_data(TensorData::new(path, [n, n]), &dev());
let rel = Tensor::from_data(TensorData::new(vec![1.0f32; d], [1, d]), &dev());
let mut ind = vec![0.0f32; n * d];
for (slot, v) in ind.iter_mut().take(d).zip(std::iter::repeat(1.0)) {
*slot = v; }
let h0 = Tensor::<B, 2>::from_data(TensorData::new(ind, [n, d]), &dev());
let mut h = Tensor::<B, 2>::zeros([n, d], &dev());
let mut prev = 0.0;
for _ in 0..4 {
h = h0.clone()
+ adj
.clone()
.matmul(h.clone() * rel.clone().slice([0..1, 0..d]));
let c = NBFConv::<B>::coverage(&h);
assert!(c > prev, "coverage must climb: {prev} -> {c}");
prev = c;
}
assert_eq!(prev, 1.0, "path fully reached: boundary + 3 hops");
}
#[test]
fn nbf_edges_matches_dense() {
let (n, d) = (4, 3);
let layer = NBFConv::<B>::init(d, &dev());
let e = [(0usize, 1usize, 0usize), (1, 2, 0), (2, 3, 1), (3, 0, 1)];
let mut a0 = vec![0.0f32; n * n];
let mut a1 = vec![0.0f32; n * n];
for &(u, v, t) in &e {
if t == 0 {
a0[v * n + u] = 1.0;
} else {
a1[v * n + u] = 1.0;
}
}
let adjs = [
Tensor::from_data(TensorData::new(a0, [n, n]), &dev()),
Tensor::from_data(TensorData::new(a1, [n, n]), &dev()),
];
let rel2 = Tensor::from_data(
TensorData::new((0..2 * d).map(|i| 0.2 + i as f32 / 9.0).collect(), [2, d]),
&dev(),
);
let h_v: Vec<f32> = (0..n * d).map(|i| i as f32 / 11.0).collect();
let h2 = Tensor::from_data(TensorData::new(h_v.clone(), [n, d]), &dev());
let h0_v: Vec<f32> = (0..n * d).map(|i| ((i * 7) % 5) as f32 / 6.0).collect();
let h02 = Tensor::from_data(TensorData::new(h0_v.clone(), [n, d]), &dev());
let dense: Vec<f32> = layer
.forward(h2, h02, &adjs, rel2.clone())
.into_data()
.to_vec()
.unwrap();
let heads = Tensor::from_data(TensorData::new(vec![0i64, 1, 2, 3], [4]), &dev());
let tails = Tensor::from_data(TensorData::new(vec![1i64, 2, 3, 0], [4]), &dev());
let etypes = Tensor::from_data(TensorData::new(vec![0i64, 0, 1, 1], [4]), &dev());
let h3 = Tensor::from_data(TensorData::new(h_v, [1, n, d]), &dev());
let h03 = Tensor::from_data(TensorData::new(h0_v, [1, n, d]), &dev());
let edges: Vec<f32> = layer
.forward_edges(h3, h03, heads, tails, etypes, rel2.reshape([1, 2, d]))
.into_data()
.to_vec()
.unwrap();
for (a, b) in dense.iter().zip(&edges) {
assert!((a - b).abs() < 1e-5, "dense {a} vs edges {b}");
}
}
#[test]
fn two_stage_conditional_propagation_composes() {
use crate::relgraph::relation_graph;
let d = 4;
let triples = [(0usize, 0usize, 1usize), (1, 1, 2), (2, 0, 3)];
let (n_ent, n_rel) = (4, 2);
let rel_nodes = 2 * n_rel;
let rgraph = relation_graph::<B>(&triples, n_rel, &dev());
let fund = Tensor::from_data(
TensorData::new((0..4 * d).map(|i| 0.1 + i as f32 / 20.0).collect(), [4, d]),
&dev(),
);
let rel_layer = NBFConv::<B>::init(d, &dev());
let query_rel = 0usize;
let mut ind = vec![0.0f32; rel_nodes * d];
for k in 0..d {
ind[query_rel * d + k] = 1.0; }
let h0r = Tensor::from_data(TensorData::new(ind, [rel_nodes, d]), &dev());
let mut hr = Tensor::zeros([rel_nodes, d], &dev());
for _ in 0..2 {
hr = rel_layer.forward(hr, h0r.clone(), &rgraph, fund.clone());
}
assert_eq!(hr.dims(), [rel_nodes, d]);
let mut adjs = Vec::new();
for r in 0..rel_nodes {
let mut m = vec![0.0f32; n_ent * n_ent];
for &(h, rr, t) in &triples {
if rr == r {
m[h * n_ent + t] = 1.0;
}
if rr + n_rel == r {
m[t * n_ent + h] = 1.0;
}
}
adjs.push(Tensor::from_data(
TensorData::new(m, [n_ent, n_ent]),
&dev(),
));
}
let ent_layer = NBFConv::<B>::init(d, &dev());
let head = 0usize;
let mut ind_e = vec![0.0f32; n_ent * d];
for k in 0..d {
ind_e[head * d + k] = 1.0;
}
let h0e = Tensor::from_data(TensorData::new(ind_e, [n_ent, d]), &dev());
let mut he = Tensor::zeros([n_ent, d], &dev());
for _ in 0..3 {
he = ent_layer.forward(he, h0e.clone(), &adjs, hr.clone());
}
assert_eq!(he.dims(), [n_ent, d]);
let mut ind2 = vec![0.0f32; rel_nodes * d];
for k in 0..d {
ind2[d + k] = 1.0;
}
let h0r2 = Tensor::from_data(TensorData::new(ind2, [rel_nodes, d]), &dev());
let mut hr2 = Tensor::zeros([rel_nodes, d], &dev());
for _ in 0..2 {
hr2 = rel_layer.forward(hr2, h0r2.clone(), &rgraph, fund.clone());
}
let mut he2 = Tensor::zeros([n_ent, d], &dev());
for _ in 0..3 {
he2 = ent_layer.forward(he2, h0e.clone(), &adjs, hr2.clone());
}
let a: Vec<f32> = he.into_data().to_vec().unwrap();
let b: Vec<f32> = he2.into_data().to_vec().unwrap();
let diff: f32 = a.iter().zip(&b).map(|(p, q)| (p - q).abs()).sum();
assert!(diff > 1e-4, "query relation must condition entity states");
}
#[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");
}
}