#![allow(dead_code)]
type Result<T, E = torsh_core::error::TorshError> = std::result::Result<T, E>;
use crate::parameter::Parameter;
use crate::{GraphData, GraphLayer};
use scirs2_core::random::thread_rng;
use torsh_tensor::{
creation::{from_vec, randn, zeros},
Tensor,
};
fn softplus(x: f32) -> f32 {
if x > 0.0 {
x + (-x).exp().ln_1p()
} else {
x.exp().ln_1p()
}
}
#[derive(Debug)]
pub struct GraphVAE {
encoder_in_features: usize,
encoder_hidden_features: usize,
latent_dim: usize,
encoder_layer1: Parameter,
encoder_layer2: Parameter,
mu_layer: Parameter,
logvar_layer: Parameter,
decoder_layer1: Parameter,
decoder_layer2: Parameter,
node_decoder: Parameter,
edge_decoder: Parameter,
beta: f32,
encoder_bias1: Option<Parameter>,
encoder_bias2: Option<Parameter>,
decoder_bias1: Option<Parameter>,
decoder_bias2: Option<Parameter>,
}
impl GraphVAE {
pub fn new(
in_features: usize,
hidden_features: usize,
latent_dim: usize,
beta: f32,
use_bias: bool,
) -> Result<Self> {
let encoder_layer1 = Parameter::new(randn(&[in_features, hidden_features])?);
let encoder_layer2 = Parameter::new(randn(&[hidden_features, hidden_features])?);
let mu_layer = Parameter::new(randn(&[hidden_features, latent_dim])?);
let logvar_layer = Parameter::new(randn(&[hidden_features, latent_dim])?);
let decoder_layer1 = Parameter::new(randn(&[latent_dim, hidden_features])?);
let decoder_layer2 = Parameter::new(randn(&[hidden_features, hidden_features])?);
let node_decoder = Parameter::new(randn(&[hidden_features, in_features])?);
let edge_decoder = Parameter::new(randn(&[hidden_features, 1])?);
let (encoder_bias1, encoder_bias2, decoder_bias1, decoder_bias2) = if use_bias {
(
Some(Parameter::new(zeros(&[hidden_features])?)),
Some(Parameter::new(zeros(&[hidden_features])?)),
Some(Parameter::new(zeros(&[hidden_features])?)),
Some(Parameter::new(zeros(&[hidden_features])?)),
)
} else {
(None, None, None, None)
};
Ok(Self {
encoder_in_features: in_features,
encoder_hidden_features: hidden_features,
latent_dim,
encoder_layer1,
encoder_layer2,
mu_layer,
logvar_layer,
decoder_layer1,
decoder_layer2,
node_decoder,
edge_decoder,
beta,
encoder_bias1,
encoder_bias2,
decoder_bias1,
decoder_bias2,
})
}
pub fn encode(&self, graph: &GraphData) -> Result<(Tensor, Tensor)> {
let mut h = graph.x.matmul(&self.encoder_layer1.clone_data())?;
if let Some(ref bias) = self.encoder_bias1 {
h = h.add(&bias.clone_data())?;
}
h = self.relu(&h)?;
h = h.matmul(&self.encoder_layer2.clone_data())?;
if let Some(ref bias) = self.encoder_bias2 {
h = h.add(&bias.clone_data())?;
}
h = self.relu(&h)?;
let graph_embedding = h.mean(Some(&[0]), false)?;
let graph_embedding_2d = graph_embedding.unsqueeze(0)?;
let mu = graph_embedding_2d.matmul(&self.mu_layer.clone_data())?;
let logvar = graph_embedding_2d.matmul(&self.logvar_layer.clone_data())?;
Ok((mu, logvar))
}
pub fn reparameterize(&self, mu: &Tensor, logvar: &Tensor) -> Result<Tensor> {
let std = logvar.mul_scalar(0.5)?.exp()?;
let epsilon = randn(mu.shape().dims())?;
Ok(mu.add(&std.mul(&epsilon)?)?)
}
pub fn decode(&self, z: &Tensor, num_nodes: usize) -> Result<GraphData> {
let mut h = z.matmul(&self.decoder_layer1.clone_data())?;
if let Some(ref bias) = self.decoder_bias1 {
h = h.add(&bias.clone_data())?;
}
h = self.relu(&h)?;
h = h.matmul(&self.decoder_layer2.clone_data())?;
if let Some(ref bias) = self.decoder_bias2 {
h = h.add(&bias.clone_data())?;
}
h = self.relu(&h)?;
let h_expanded = self.expand_to_nodes(&h, num_nodes)?;
let node_features = h_expanded.matmul(&self.node_decoder.clone_data())?;
let edge_logits = self.decode_edges(&h_expanded, num_nodes)?;
let edge_index = self.sample_edges(&edge_logits, num_nodes)?;
Ok(GraphData::new(node_features, edge_index))
}
pub fn forward(&self, graph: &GraphData) -> Result<(GraphData, Tensor, Tensor)> {
let (mu, logvar) = self.encode(graph)?;
let z = self.reparameterize(&mu, &logvar)?;
let reconstructed = self.decode(&z, graph.num_nodes)?;
Ok((reconstructed, mu, logvar))
}
pub fn compute_loss(
&self,
graph: &GraphData,
reconstructed: &GraphData,
mu: &Tensor,
logvar: &Tensor,
) -> Result<f32> {
let recon_loss = self.reconstruction_loss(graph, reconstructed)?;
let kl_loss = self.kl_divergence(mu, logvar)?;
Ok(recon_loss + self.beta * kl_loss)
}
fn reconstruction_loss(&self, original: &GraphData, reconstructed: &GraphData) -> Result<f32> {
let orig_data = original.x.to_vec()?;
let recon_data = reconstructed.x.to_vec()?;
let mut mse = 0.0;
let len = orig_data.len().min(recon_data.len());
for i in 0..len {
mse += (orig_data[i] - recon_data[i]).powi(2);
}
Ok(mse / len as f32)
}
fn kl_divergence(&self, mu: &Tensor, logvar: &Tensor) -> Result<f32> {
let mu_data = mu.to_vec()?;
let logvar_data = logvar.to_vec()?;
let mut kl = 0.0;
for i in 0..mu_data.len() {
kl += -0.5 * (1.0 + logvar_data[i] - mu_data[i].powi(2) - logvar_data[i].exp());
}
Ok(kl / mu_data.len() as f32)
}
pub fn generate(&self, num_nodes: usize) -> Result<GraphData> {
let z = randn(&[1, self.latent_dim])?;
self.decode(&z, num_nodes)
}
pub fn interpolate(
&self,
graph1: &GraphData,
graph2: &GraphData,
alpha: f32,
num_nodes: usize,
) -> Result<GraphData> {
let (mu1, _) = self.encode(graph1)?;
let (mu2, _) = self.encode(graph2)?;
let z_interp = mu1.mul_scalar(1.0 - alpha)?.add(&mu2.mul_scalar(alpha)?)?;
self.decode(&z_interp, num_nodes)
}
fn relu(&self, x: &Tensor) -> Result<Tensor> {
let data = x.to_vec()?;
let activated: Vec<f32> = data.iter().map(|&v| v.max(0.0)).collect();
Ok(from_vec(
activated,
x.shape().dims(),
torsh_core::device::DeviceType::Cpu,
)?)
}
fn expand_to_nodes(&self, h: &Tensor, num_nodes: usize) -> Result<Tensor> {
let h_data = h.to_vec()?;
let feat_dim = h_data.len();
let mut expanded_data = Vec::new();
for _ in 0..num_nodes {
expanded_data.extend(&h_data);
}
Ok(from_vec(
expanded_data,
&[num_nodes, feat_dim],
torsh_core::device::DeviceType::Cpu,
)?)
}
fn decode_edges(&self, h: &Tensor, num_nodes: usize) -> Result<Tensor> {
let mut edge_logits_data = Vec::new();
for i in 0..num_nodes {
for j in 0..num_nodes {
if i != j {
let h_i = h.slice_tensor(0, i, i + 1)?;
let h_j = h.slice_tensor(0, j, j + 1)?;
let logit = h_i.dot(&h_j.t()?)?.item()?;
edge_logits_data.push(logit);
} else {
edge_logits_data.push(-1000.0); }
}
}
Ok(from_vec(
edge_logits_data,
&[num_nodes, num_nodes],
torsh_core::device::DeviceType::Cpu,
)?)
}
fn sample_edges(&self, edge_logits: &Tensor, num_nodes: usize) -> Result<Tensor> {
let logits_data = edge_logits.to_vec()?;
let mut edges = Vec::new();
for i in 0..num_nodes {
for j in 0..num_nodes {
if i != j {
let idx = i * num_nodes + j;
let prob = 1.0 / (1.0 + (-logits_data[idx]).exp());
if prob > 0.5 {
edges.push(i as f32);
edges.push(j as f32);
}
}
}
}
if edges.is_empty() {
return Ok(zeros(&[2, 0])?);
}
let num_edges = edges.len() / 2;
Ok(from_vec(
edges,
&[2, num_edges],
torsh_core::device::DeviceType::Cpu,
)?)
}
}
impl GraphLayer for GraphVAE {
fn forward(&self, graph: &GraphData) -> Result<GraphData> {
let (reconstructed, _, _) = GraphVAE::forward(self, graph)?;
Ok(reconstructed)
}
fn parameters(&self) -> Vec<Tensor> {
let mut params = vec![
self.encoder_layer1.clone_data(),
self.encoder_layer2.clone_data(),
self.mu_layer.clone_data(),
self.logvar_layer.clone_data(),
self.decoder_layer1.clone_data(),
self.decoder_layer2.clone_data(),
self.node_decoder.clone_data(),
self.edge_decoder.clone_data(),
];
if let Some(ref b) = self.encoder_bias1 {
params.push(b.clone_data());
}
if let Some(ref b) = self.encoder_bias2 {
params.push(b.clone_data());
}
if let Some(ref b) = self.decoder_bias1 {
params.push(b.clone_data());
}
if let Some(ref b) = self.decoder_bias2 {
params.push(b.clone_data());
}
params
}
}
#[derive(Debug)]
pub struct GraphGAN {
latent_dim: usize,
hidden_dim: usize,
output_features: usize,
generator: GraphGANGenerator,
discriminator: GraphGANDiscriminator,
}
impl GraphGAN {
pub fn new(
latent_dim: usize,
hidden_dim: usize,
output_features: usize,
use_bias: bool,
) -> Result<Self> {
let generator = GraphGANGenerator::new(latent_dim, hidden_dim, output_features, use_bias)?;
let discriminator = GraphGANDiscriminator::new(output_features, hidden_dim, use_bias)?;
Ok(Self {
latent_dim,
hidden_dim,
output_features,
generator,
discriminator,
})
}
pub fn generate(&self, num_nodes: usize) -> Result<GraphData> {
let z = randn(&[1, self.latent_dim])?;
self.generator.generate(&z, num_nodes)
}
pub fn discriminate(&self, graph: &GraphData) -> Result<f32> {
self.discriminator.forward(graph)
}
pub fn discriminate_logit(&self, graph: &GraphData) -> Result<f32> {
self.discriminator.forward_logit(graph)
}
pub fn generator_loss(&self, num_nodes: usize) -> Result<f32> {
let fake_graph = self.generate(num_nodes)?;
let fake_logit = self.discriminate_logit(&fake_graph)?;
Ok(softplus(-fake_logit))
}
pub fn discriminator_loss(&self, real_graph: &GraphData, num_nodes: usize) -> Result<f32> {
let real_logit = self.discriminate_logit(real_graph)?;
let fake_graph = self.generate(num_nodes)?;
let fake_logit = self.discriminate_logit(&fake_graph)?;
Ok(softplus(-real_logit) + softplus(fake_logit))
}
pub fn generator_parameters(&self) -> Vec<Tensor> {
self.generator.parameters()
}
pub fn discriminator_parameters(&self) -> Vec<Tensor> {
self.discriminator.parameters()
}
}
#[derive(Debug)]
struct GraphGANGenerator {
latent_dim: usize,
hidden_dim: usize,
output_features: usize,
layer1: Parameter,
layer2: Parameter,
node_layer: Parameter,
edge_layer: Parameter,
bias1: Option<Parameter>,
bias2: Option<Parameter>,
}
impl GraphGANGenerator {
fn new(
latent_dim: usize,
hidden_dim: usize,
output_features: usize,
use_bias: bool,
) -> Result<Self> {
let layer1 = Parameter::new(randn(&[latent_dim, hidden_dim])?);
let layer2 = Parameter::new(randn(&[hidden_dim, hidden_dim])?);
let node_layer = Parameter::new(randn(&[hidden_dim, output_features])?);
let edge_layer = Parameter::new(randn(&[hidden_dim, 1])?);
let (bias1, bias2) = if use_bias {
(
Some(Parameter::new(zeros(&[hidden_dim])?)),
Some(Parameter::new(zeros(&[hidden_dim])?)),
)
} else {
(None, None)
};
Ok(Self {
latent_dim,
hidden_dim,
output_features,
layer1,
layer2,
node_layer,
edge_layer,
bias1,
bias2,
})
}
fn generate(&self, z: &Tensor, num_nodes: usize) -> Result<GraphData> {
let mut h = z.matmul(&self.layer1.clone_data())?;
if let Some(ref bias) = self.bias1 {
h = h.add(&bias.clone_data())?;
}
h = self.leaky_relu(&h, 0.2)?;
h = h.matmul(&self.layer2.clone_data())?;
if let Some(ref bias) = self.bias2 {
h = h.add(&bias.clone_data())?;
}
h = self.leaky_relu(&h, 0.2)?;
let h_expanded = self.expand_to_nodes(&h, num_nodes)?;
let node_features = h_expanded.matmul(&self.node_layer.clone_data())?;
let node_features = self.tanh(&node_features)?;
let edge_index = self.generate_edges(&h_expanded, num_nodes)?;
Ok(GraphData::new(node_features, edge_index))
}
fn leaky_relu(&self, x: &Tensor, alpha: f32) -> Result<Tensor> {
let data = x.to_vec()?;
let activated: Vec<f32> = data
.iter()
.map(|&v| if v > 0.0 { v } else { alpha * v })
.collect();
Ok(from_vec(
activated,
x.shape().dims(),
torsh_core::device::DeviceType::Cpu,
)?)
}
fn tanh(&self, x: &Tensor) -> Result<Tensor> {
let data = x.to_vec()?;
let activated: Vec<f32> = data.iter().map(|&v| v.tanh()).collect();
Ok(from_vec(
activated,
x.shape().dims(),
torsh_core::device::DeviceType::Cpu,
)?)
}
fn expand_to_nodes(&self, h: &Tensor, num_nodes: usize) -> Result<Tensor> {
let h_data = h.to_vec()?;
let feat_dim = h_data.len();
let mut expanded_data = Vec::new();
for _ in 0..num_nodes {
expanded_data.extend(&h_data);
}
Ok(from_vec(
expanded_data,
&[num_nodes, feat_dim],
torsh_core::device::DeviceType::Cpu,
)?)
}
fn generate_edges(&self, _h: &Tensor, num_nodes: usize) -> Result<Tensor> {
let mut edges = Vec::new();
let mut rng = thread_rng();
for i in 0..num_nodes {
for j in (i + 1)..num_nodes {
if rng.gen_range(0.0..1.0) > 0.7 {
edges.push(i as f32);
edges.push(j as f32);
edges.push(j as f32);
edges.push(i as f32);
}
}
}
if edges.is_empty() {
return Ok(zeros(&[2, 0])?);
}
let num_edges = edges.len() / 2;
Ok(from_vec(
edges,
&[2, num_edges],
torsh_core::device::DeviceType::Cpu,
)?)
}
fn parameters(&self) -> Vec<Tensor> {
let mut params = vec![
self.layer1.clone_data(),
self.layer2.clone_data(),
self.node_layer.clone_data(),
self.edge_layer.clone_data(),
];
if let Some(ref b) = self.bias1 {
params.push(b.clone_data());
}
if let Some(ref b) = self.bias2 {
params.push(b.clone_data());
}
params
}
}
#[derive(Debug)]
struct GraphGANDiscriminator {
input_features: usize,
hidden_dim: usize,
layer1: Parameter,
layer2: Parameter,
output_layer: Parameter,
bias1: Option<Parameter>,
bias2: Option<Parameter>,
bias_out: Option<Parameter>,
}
impl GraphGANDiscriminator {
fn new(input_features: usize, hidden_dim: usize, use_bias: bool) -> Result<Self> {
let layer1 = Parameter::new(randn(&[input_features, hidden_dim])?);
let layer2 = Parameter::new(randn(&[hidden_dim, hidden_dim])?);
let output_layer = Parameter::new(randn(&[hidden_dim, 1])?);
let (bias1, bias2, bias_out) = if use_bias {
(
Some(Parameter::new(zeros(&[hidden_dim])?)),
Some(Parameter::new(zeros(&[hidden_dim])?)),
Some(Parameter::new(zeros(&[1])?)),
)
} else {
(None, None, None)
};
Ok(Self {
input_features,
hidden_dim,
layer1,
layer2,
output_layer,
bias1,
bias2,
bias_out,
})
}
fn forward_logit(&self, graph: &GraphData) -> Result<f32> {
let mut h = graph.x.matmul(&self.layer1.clone_data())?;
if let Some(ref bias) = self.bias1 {
h = h.add(&bias.clone_data())?;
}
h = self.leaky_relu(&h, 0.2)?;
h = h.matmul(&self.layer2.clone_data())?;
if let Some(ref bias) = self.bias2 {
h = h.add(&bias.clone_data())?;
}
h = self.leaky_relu(&h, 0.2)?;
let graph_repr = h.mean(Some(&[0]), false)?;
let graph_repr_2d = graph_repr.unsqueeze(0)?;
let mut logit = graph_repr_2d.matmul(&self.output_layer.clone_data())?;
if let Some(ref bias) = self.bias_out {
logit = logit.add(&bias.clone_data())?;
}
logit.item()
}
fn forward(&self, graph: &GraphData) -> Result<f32> {
let logit_val = self.forward_logit(graph)?;
Ok(1.0 / (1.0 + (-logit_val).exp()))
}
fn leaky_relu(&self, x: &Tensor, alpha: f32) -> Result<Tensor> {
let data = x.to_vec()?;
let activated: Vec<f32> = data
.iter()
.map(|&v| if v > 0.0 { v } else { alpha * v })
.collect();
Ok(from_vec(
activated,
x.shape().dims(),
torsh_core::device::DeviceType::Cpu,
)?)
}
fn parameters(&self) -> Vec<Tensor> {
let mut params = vec![
self.layer1.clone_data(),
self.layer2.clone_data(),
self.output_layer.clone_data(),
];
if let Some(ref b) = self.bias1 {
params.push(b.clone_data());
}
if let Some(ref b) = self.bias2 {
params.push(b.clone_data());
}
if let Some(ref b) = self.bias_out {
params.push(b.clone_data());
}
params
}
}
#[derive(Debug)]
pub struct ConditionalGraphGenerator {
vae: GraphVAE,
condition_dim: usize,
condition_layer: Parameter,
}
impl ConditionalGraphGenerator {
pub fn new(
in_features: usize,
hidden_features: usize,
latent_dim: usize,
condition_dim: usize,
beta: f32,
) -> Result<Self> {
let vae = GraphVAE::new(in_features, hidden_features, latent_dim, beta, true)?;
let condition_layer = Parameter::new(randn(&[condition_dim, latent_dim])?);
Ok(Self {
vae,
condition_dim,
condition_layer,
})
}
pub fn generate_conditional(&self, condition: &Tensor, num_nodes: usize) -> Result<GraphData> {
let condition_bias = condition.matmul(&self.condition_layer.clone_data())?;
let z_base = randn(&[1, self.vae.latent_dim])?;
let z = z_base.add(&condition_bias)?;
self.vae.decode(&z, num_nodes)
}
fn parameters(&self) -> Vec<Tensor> {
let mut params = self.vae.parameters();
params.push(self.condition_layer.clone_data());
params
}
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::DeviceType;
#[test]
fn test_graphvae_creation() {
let vae = GraphVAE::new(8, 16, 10, 1.0, true).expect("operation should succeed");
assert_eq!(vae.encoder_in_features, 8);
assert_eq!(vae.encoder_hidden_features, 16);
assert_eq!(vae.latent_dim, 10);
assert_eq!(vae.beta, 1.0);
}
#[test]
fn test_graphvae_encode_decode() {
let features = randn(&[5, 8]).unwrap();
let edges = vec![0.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 4.0];
let edge_index = from_vec(edges, &[2, 4], DeviceType::Cpu).unwrap();
let graph = GraphData::new(features, edge_index);
let vae = GraphVAE::new(8, 16, 10, 1.0, true).expect("operation should succeed");
let (mu, logvar) = vae.encode(&graph).expect("operation should succeed");
assert_eq!(mu.shape().dims(), &[1, 10]);
assert_eq!(logvar.shape().dims(), &[1, 10]);
let z = vae
.reparameterize(&mu, &logvar)
.expect("operation should succeed");
assert_eq!(z.shape().dims(), &[1, 10]);
let reconstructed = vae.decode(&z, 5).expect("operation should succeed");
assert_eq!(reconstructed.num_nodes, 5);
}
#[test]
fn test_graphvae_generation() {
let vae = GraphVAE::new(8, 16, 10, 1.0, true).expect("operation should succeed");
let generated = vae.generate(6).expect("operation should succeed");
assert_eq!(generated.num_nodes, 6);
assert_eq!(generated.x.shape().dims()[0], 6);
assert_eq!(generated.x.shape().dims()[1], 8);
}
#[test]
fn test_graphvae_interpolation() {
let features1 = randn(&[4, 6]).unwrap();
let features2 = randn(&[4, 6]).unwrap();
let edges = vec![0.0, 1.0, 1.0, 2.0, 2.0, 3.0];
let edge_index = from_vec(edges.clone(), &[2, 3], DeviceType::Cpu).unwrap();
let graph1 = GraphData::new(features1, edge_index.clone());
let graph2 = GraphData::new(features2, edge_index);
let vae = GraphVAE::new(6, 12, 8, 1.0, true).expect("operation should succeed");
let interpolated = vae
.interpolate(&graph1, &graph2, 0.5, 4)
.expect("operation should succeed");
assert_eq!(interpolated.num_nodes, 4);
}
#[test]
fn test_graphgan_creation() {
let gan = GraphGAN::new(16, 32, 8, true).expect("operation should succeed");
assert_eq!(gan.latent_dim, 16);
assert_eq!(gan.hidden_dim, 32);
assert_eq!(gan.output_features, 8);
}
#[test]
fn test_graphgan_generation() {
let gan = GraphGAN::new(16, 32, 8, true).expect("operation should succeed");
let generated = gan.generate(5).expect("operation should succeed");
assert_eq!(generated.num_nodes, 5);
assert_eq!(generated.x.shape().dims()[1], 8);
}
#[test]
fn test_graphgan_discriminate() {
let features = randn(&[4, 8]).unwrap();
let edges = vec![0.0, 1.0, 1.0, 2.0, 2.0, 3.0];
let edge_index = from_vec(edges, &[2, 3], DeviceType::Cpu).unwrap();
let graph = GraphData::new(features, edge_index);
let gan = GraphGAN::new(16, 32, 8, true).expect("operation should succeed");
let score = gan.discriminate(&graph).expect("operation should succeed");
assert!(score >= 0.0 && score <= 1.0);
}
#[test]
fn test_conditional_generation() {
let cond_gen =
ConditionalGraphGenerator::new(8, 16, 10, 4, 1.0).expect("operation should succeed");
let condition = randn(&[1, 4]).unwrap();
let generated = cond_gen
.generate_conditional(&condition, 5)
.expect("operation should succeed");
assert_eq!(generated.num_nodes, 5);
assert_eq!(generated.x.shape().dims()[1], 8);
}
#[test]
fn test_graphvae_loss_computation() {
let features = randn(&[3, 6]).unwrap();
let edges = vec![0.0, 1.0, 1.0, 2.0];
let edge_index = from_vec(edges, &[2, 2], DeviceType::Cpu).unwrap();
let graph = GraphData::new(features, edge_index);
let vae = GraphVAE::new(6, 12, 8, 1.0, true).expect("operation should succeed");
let (reconstructed, mu, logvar) = vae.forward(&graph).expect("operation should succeed");
let loss = vae
.compute_loss(&graph, &reconstructed, &mu, &logvar)
.expect("operation should succeed");
assert!(loss > 0.0);
}
fn deterministic_unit_interval(name: &str, index: u64) -> f32 {
let mut hash: u64 = 0xcbf2_9ce4_8422_2325; for byte in name.bytes() {
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3); }
hash ^= index.wrapping_add(0x9E37_79B9_7F4A_7C15);
hash = (hash ^ (hash >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
hash = (hash ^ (hash >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
hash ^= hash >> 31;
((hash >> 40) as f32) / (1u64 << 24) as f32
}
fn reinit_params(params: &[Tensor], label: &str) {
for (i, tensor) in params.iter().enumerate() {
let dims = tensor.shape().dims().to_vec();
if dims.len() != 2 {
continue;
}
let numel: usize = dims.iter().product();
let bound = (6.0 / (dims[0] + dims[1]) as f32).sqrt();
let name = format!("{label}.param{i}");
let values: Vec<f32> = (0..numel)
.map(|j| bound * (2.0 * deterministic_unit_interval(&name, j as u64) - 1.0))
.collect();
tensor
.set_slice(0, &values)
.expect("deterministic reinit set_slice should succeed");
}
}
fn deterministic_reinit_gan(gan: &GraphGAN) {
reinit_params(&gan.generator_parameters(), "gan.generator");
reinit_params(&gan.discriminator_parameters(), "gan.discriminator");
}
#[test]
fn test_graphgan_losses() {
let features = randn(&[4, 8]).unwrap();
let edges = vec![0.0, 1.0, 1.0, 2.0, 2.0, 3.0];
let edge_index = from_vec(edges, &[2, 3], DeviceType::Cpu).unwrap();
let graph = GraphData::new(features, edge_index);
let gan = GraphGAN::new(16, 32, 8, true).expect("operation should succeed");
deterministic_reinit_gan(&gan);
let gen_loss = gan.generator_loss(4).expect("operation should succeed");
assert!(gen_loss > 0.0);
let disc_loss = gan
.discriminator_loss(&graph, 4)
.expect("operation should succeed");
assert!(disc_loss.is_finite());
}
#[test]
fn discriminator_logit_stays_bounded_across_many_random_graphs() {
let gan = GraphGAN::new(16, 32, 8, true).expect("operation should succeed");
deterministic_reinit_gan(&gan);
for _ in 0..2000 {
let features = randn(&[4, 8]).expect("randn");
let edges = vec![0.0, 1.0, 1.0, 2.0, 2.0, 3.0];
let edge_index = from_vec(edges, &[2, 3], DeviceType::Cpu).expect("from_vec");
let graph = GraphData::new(features, edge_index);
let real_logit = gan.discriminate_logit(&graph).expect("discriminate_logit");
assert!(
real_logit.abs() < 50.0,
"discriminator logit strayed far enough from 0 to approach the \
softplus underflow edge: {real_logit}"
);
let gen_loss = gan.generator_loss(4).expect("generator_loss");
assert!(gen_loss > 0.0, "gen_loss was not > 0.0: {gen_loss}");
let disc_loss = gan
.discriminator_loss(&graph, 4)
.expect("discriminator_loss");
assert!(
disc_loss.is_finite(),
"disc_loss was non-finite: {disc_loss}"
);
}
}
}