1
Fork 0

feat: more generic naming for critic, since messager will use the same

This commit is contained in:
cedric 2026-04-02 18:11:49 +00:00
parent 729f122be4
commit 40c7e32344
2 changed files with 4 additions and 3 deletions

View file

@ -21,7 +21,7 @@ class SemiGenericNetwork(nn.Module):
return x
class Critic(nn.Module):
class OneDenseLayerMLP(nn.Module):
@nn.compact
def __call__(self, x):
return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x)

View file

@ -19,7 +19,7 @@ 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, Critic, AgentParams, Storage
from MLPs.mlps import SemiGenericNetwork, Actor, OneDenseLayerMLP, AgentParams, Storage
from ppo import PPO
@ -119,7 +119,8 @@ def train(args: PPOArgs):
network = SemiGenericNetwork()
critic_network = SemiGenericNetwork()
actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX
critic = Critic()
critic = OneDenseLayerMLP()
# messager = OneDenseLayerMLP()
sample_obs = jnp.concatenate(
[