feat: adjacency builder
This commit is contained in:
parent
8830289992
commit
2777d4f6da
2 changed files with 149 additions and 13 deletions
|
|
@ -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
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