1
Fork 0

feat: adjacency builder

This commit is contained in:
Cedric 2026-04-28 06:46:26 +00:00
parent 8830289992
commit 2777d4f6da
2 changed files with 149 additions and 13 deletions

View file

@ -27,10 +27,11 @@ from brittle_star_project.ppo import PPO
from enum import Enum
class ObsMode(Enum):
CENTRALIZED = "centralized"
ARM = "arm"
SEGMENT = "segment"
class MorphMode(Enum):
CENTRALIZED = 0
FULLY_CONNECTED = 1
RING = 2
SEGMENT = 3
# TODO: move to config
@ -49,6 +50,71 @@ _ALLOWED_OBS_KEYS = {
# TODO: clip scaled reward?
def build_adjacency(segments_per_arm, mode: MorphMode):
num_arms = len(segments_per_arm)
num_segments = sum(segments_per_arm)
# FOR NOW SEMI HARDCODE:
# CENTRALIZED: 1 agent, no stress, adja = 1,1 = [[1]]
# FULLY CONNECTED: 5 agents: adj = alle 1
# CENTRAL DISK:#arms= 5 agents, only neighbor as adjacent so diagonal kinda..
# ARM = #segments agents: diago kinda, but extra, center ring too, put center mlps first or..
if mode == MorphMode.CENTRALIZED:
return jnp.ones((1, 1))
if mode == MorphMode.FULLY_CONNECTED:
adj = jnp.ones((num_arms, num_arms)) # everybody adjacent everybody
return adj
if mode == MorphMode.RING: # ring
adj = jnp.zeros((num_arms, num_arms))
for i in range(num_arms):
adj = adj.at[i, i].set(1) # self
adj = adj.at[i, (i - 1) % num_arms].set(1)
adj = adj.at[i, (i + 1) % num_arms].set(1) # left and right..
return adj
if mode == MorphMode.SEGMENT:
num_nodes = num_arms + num_segments
adj = jnp.zeros((num_nodes, num_nodes))
# first ring
for i in range(num_arms):
# self
adj = adj.at[i, i].set(1)
# ring neighbors
adj = adj.at[i, (i - 1) % num_arms].set(1)
adj = adj.at[i, (i + 1) % num_arms].set(1)
# then segment chains
idx = 0
for arm_idx, seg_count in enumerate(segments_per_arm):
for i in range(seg_count):
seg_node = num_arms + idx + i
adj = adj.at[seg_node, seg_node].set(1)
if i > 0:
adj = adj.at[seg_node, seg_node - 1].set(1)
if i < seg_count - 1:
adj = adj.at[seg_node, seg_node + 1].set(1)
idx += seg_count
idx = 0
for arm_idx, seg_count in enumerate(segments_per_arm):
first_seg = num_arms + idx # first segment of this arm
# connect ring node first segment
adj = adj.at[arm_idx, first_seg].set(1)
adj = adj.at[first_seg, arm_idx].set(1)
idx += seg_count
return adj
@jax.jit
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
return jnp.clip(action, low, high)
@ -71,15 +137,6 @@ def _normalize_obs(obs, mean, var, eps=1e-8):
return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0)
from enum import Enum
class ObsMode(Enum):
CENTRALIZED = 0
ARM = 1
SEGMENT = 2
@jax.jit
def _convert_obs_dict_to_array(obs_dict, obs_mode, segments_per_arm):

79
tests/test_adjacency.py Normal file
View 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)