From 6604433aa78b34bd3bfed68977bc75b8febf8df6 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 17 Apr 2026 10:09:19 +0200 Subject: [PATCH] backup --- docs/design/network_inputs.md | 19 +++++++++++++++++++ src/brittle_star_project/MLPs/mlps.py | 15 +++++++++++---- .../trainers/PPOTrainer.py | 10 ++++++++-- 3 files changed, 38 insertions(+), 6 deletions(-) create mode 100644 docs/design/network_inputs.md diff --git a/docs/design/network_inputs.md b/docs/design/network_inputs.md new file mode 100644 index 0000000..eba3df0 --- /dev/null +++ b/docs/design/network_inputs.md @@ -0,0 +1,19 @@ +# Neural network architecture and inputs + +## Inputs +From the observations, we use: +* joint_position +* joint_velocity +* joint_actuator_force +* actuator_force +* disk_position +* disk_rotation +* disk_linear_velocity +* disk_angular_velocity +* unit_xy_direction_to_target +* xy_distance_to_target + +as inputs for the sensor and the feature extractor. + +## Sensor architecture +todo \ No newline at end of file diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 83d7389..1911f8c 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -2,7 +2,6 @@ from dataclasses import dataclass, fields, field import flax import flax.linen as nn -import jax.numpy as jnp import jax.tree_util from typing import Sequence, Callable from flax.linen.initializers import constant, orthogonal @@ -13,18 +12,26 @@ class GenericDenseLayersWithActivation(nn.Module): layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) activation: Callable = nn.tanh + def setup(self): + self.dense_layers = [nn.Dense(size) for size in self.layer_sizes] + @nn.compact def __call__(self, x): - for size in self.layer_sizes: - x = nn.Dense(size, kernel_init=orthogonal(jnp.sqrt(2)))(x) + for layer in self.dense_layers: + x = layer(x) x = self.activation(x) return x class OneDenseLayerMLP(nn.Module): + feature_dim: int = field(default=1) + + def setup(self): + self.dense_layer = nn.Dense(self.feature_dim) + @nn.compact def __call__(self, x): - return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) + return self.dense_layer(x) class Actor(nn.Module): diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 197036b..51f8a2b 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -1,4 +1,5 @@ import datetime +import logging import random import time from dataclasses import asdict, dataclass @@ -215,6 +216,9 @@ class PPOTrainer: self.run_name = run_name self.logger = get_logger() + # TODO: remove + self.logger.set_level(logging.DEBUG) + self.key = jax.random.PRNGKey(args.seed) self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent() @@ -262,8 +266,10 @@ class PPOTrainer: def _init_agent(self): self.logger.info("[AGENT]: Initializing agent...") - sensor = GenericDenseLayersWithActivation() - feature_extractor = GenericDenseLayersWithActivation() + sizes = [195 // 2, 195 // 4, 195 // 4] + + sensor = GenericDenseLayersWithActivation(layer_sizes=sizes) + feature_extractor = GenericDenseLayersWithActivation(layer_sizes=sizes) actor = Actor( action_dim=self.env.single_action_space.shape[0] ) # continuous actions for MJX