1
Fork 0

fix: uniform generic network naming

This commit is contained in:
cedric 2026-04-02 18:18:10 +00:00
parent 40c7e32344
commit 79221372df
2 changed files with 10 additions and 4 deletions

View file

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

View file

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