diff --git a/src/MLPs/fully_arm.py b/src/MLPs/fully_arm.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/MLPs/helpers.py b/src/MLPs/helpers.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/MLPs/centralized.py b/src/MLPs/mlps.py similarity index 50% rename from src/MLPs/centralized.py rename to src/MLPs/mlps.py index 14e3a22..f97452d 100644 --- a/src/MLPs/centralized.py +++ b/src/MLPs/mlps.py @@ -1,19 +1,20 @@ -from typing import Sequence +from typing import Sequence, Callable import flax.linen as nn import jax.numpy as jnp from flax.linen.initializers import constant, orthogonal -class Network(nn.Module): - hidden_size: int = 256 +# example usage: network = SemiGenericNetwork(layer_sizes=[256, 256], activation=nn.relu) +# semi generic so we can easily make a config for it in experiments +class SemiGenericNetwork(nn.Module): + layer_sizes: Sequence[int] = [64, 64] # default 2 layers of 64 neurons + activation: Callable = nn.tanh # default tanh @nn.compact def __call__(self, x): - # x shape: (batch, obs_dim) - x = nn.Dense(self.hidden_size, kernel_init=orthogonal(jnp.sqrt(2)))(x) - x = nn.tanh(x) - x = nn.Dense(self.hidden_size, kernel_init=orthogonal(jnp.sqrt(2)))(x) - x = nn.tanh(x) + for size in self.layer_sizes: + x = nn.Dense(size, kernel_init=orthogonal(jnp.sqrt(2)))(x) + x = self.activation(x) return x diff --git a/src/MLPs/ring_arm.py b/src/MLPs/ring_arm.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/MLPs/segment.py b/src/MLPs/segment.py deleted file mode 100644 index e69de29..0000000