feat: adjacency builder
This commit is contained in:
parent
8830289992
commit
2777d4f6da
2 changed files with 149 additions and 13 deletions
79
tests/test_adjacency.py
Normal file
79
tests/test_adjacency.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
import jax.numpy as jnp
|
||||
import pytest
|
||||
|
||||
from brittle_star_project.trainers.PPOTrainer import build_adjacency, MorphMode
|
||||
|
||||
|
||||
def assert_symmetric(adj):
|
||||
assert jnp.all(adj == adj.T)
|
||||
|
||||
|
||||
def test_centralized():
|
||||
adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.CENTRALIZED)
|
||||
|
||||
assert adj.shape == (1, 1)
|
||||
assert adj[0, 0] == 1
|
||||
|
||||
|
||||
def test_fully_connected():
|
||||
adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.FULLY_CONNECTED)
|
||||
|
||||
assert adj.shape == (5, 5)
|
||||
assert jnp.all(adj == 1)
|
||||
|
||||
|
||||
def test_ring():
|
||||
adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.RING)
|
||||
|
||||
assert adj.shape == (5, 5)
|
||||
assert_symmetric(adj)
|
||||
|
||||
# each node should connect to itself + 2 neighbors
|
||||
for i in range(5):
|
||||
assert adj[i, i] == 1
|
||||
assert jnp.sum(adj[i]) == 3
|
||||
|
||||
|
||||
def test_segment_structure():
|
||||
segments = [4, 4, 4, 4, 4]
|
||||
adj = build_adjacency(segments, MorphMode.SEGMENT)
|
||||
|
||||
num_arms = 5
|
||||
num_segments = sum(segments)
|
||||
num_nodes = num_arms + num_segments
|
||||
|
||||
assert adj.shape == (num_nodes, num_nodes)
|
||||
|
||||
# --- ring connectivity ---
|
||||
for i in range(num_arms):
|
||||
assert adj[i, i] == 1
|
||||
assert adj[i, (i - 1) % num_arms] == 1
|
||||
assert adj[i, (i + 1) % num_arms] == 1
|
||||
|
||||
# --- segment chain checks ---
|
||||
offset = num_arms
|
||||
for arm in range(5):
|
||||
for i in range(4):
|
||||
node = offset + arm * 4 + i
|
||||
|
||||
# self
|
||||
assert adj[node, node] == 1
|
||||
|
||||
# chain neighbors
|
||||
if i > 0:
|
||||
assert adj[node, node - 1] == 1
|
||||
if i < 3:
|
||||
assert adj[node, node + 1] == 1
|
||||
|
||||
# --- ring ↔ segment connections ---
|
||||
for arm in range(5):
|
||||
first_seg = num_arms + arm * 4
|
||||
assert adj[arm, first_seg] == 1
|
||||
assert adj[first_seg, arm] == 1
|
||||
|
||||
|
||||
def test_no_isolated_nodes():
|
||||
adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.SEGMENT)
|
||||
|
||||
# no node should be completely isolated
|
||||
assert jnp.all(jnp.sum(adj, axis=0) > 0)
|
||||
Reference in a new issue