1
Fork 0
This commit is contained in:
Robin Meersman 2026-04-17 10:09:19 +02:00
parent 0ed2dfe4ef
commit 6604433aa7
3 changed files with 38 additions and 6 deletions

View file

@ -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

View file

@ -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):

View file

@ -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