From 2777d4f6daf48d7c8be0696251397194c9f05a0d Mon Sep 17 00:00:00 2001 From: Cedric Date: Tue, 28 Apr 2026 06:46:26 +0000 Subject: [PATCH] feat: adjacency builder --- .../trainers/PPOTrainer.py | 83 ++++++++++++++++--- tests/test_adjacency.py | 79 ++++++++++++++++++ 2 files changed, 149 insertions(+), 13 deletions(-) create mode 100644 tests/test_adjacency.py diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index df2c571..2642446 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -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): diff --git a/tests/test_adjacency.py b/tests/test_adjacency.py new file mode 100644 index 0000000..5c33038 --- /dev/null +++ b/tests/test_adjacency.py @@ -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)