From 79221372dfca74b10ce74326900a7a963e059013 Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 18:18:10 +0000 Subject: [PATCH] fix: uniform generic network naming --- src/MLPs/mlps.py | 2 +- src/train.py | 12 +++++++++--- 2 files changed, 10 insertions(+), 4 deletions(-) diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index 1bae0c4..1fdc6a2 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -9,7 +9,7 @@ from flax.linen.initializers import constant, orthogonal # semi generic so we can easily make a config for it in experiments -class SemiGenericNetwork(nn.Module): +class GenericDenseLayersWithActivation(nn.Module): layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) activation: Callable = nn.tanh diff --git a/src/train.py b/src/train.py index 7528f26..10f14bf 100644 --- a/src/train.py +++ b/src/train.py @@ -19,7 +19,13 @@ from torch.utils.tensorboard import SummaryWriter from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from MLPs.mlps import SemiGenericNetwork, Actor, OneDenseLayerMLP, AgentParams, Storage +from MLPs.mlps import ( + GenericDenseLayersWithActivation, + Actor, + OneDenseLayerMLP, + AgentParams, + Storage, +) from ppo import PPO @@ -116,8 +122,8 @@ def train(args: PPOArgs): return args.learning_rate * frac print("Initializing the models...") - network = SemiGenericNetwork() - critic_network = SemiGenericNetwork() + network = GenericDenseLayersWithActivation() + critic_network = GenericDenseLayersWithActivation() actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX critic = OneDenseLayerMLP() # messager = OneDenseLayerMLP()