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()