class GaussianHead(nnx.Module):
"""Mean network plus a state-independent learned log standard deviation."""
def __init__(self, obs_dim, act_dim, hidden, rngs):
self.mean = nnx.Sequential(nnx.Linear(obs_dim, hidden, rngs=rngs),
jnp.tanh,
nnx.Linear(hidden, act_dim, rngs=rngs))
self.log_std = nnx.Param(jnp.zeros(act_dim))
def __call__(self, obs):
return self.mean(obs), jnp.exp(self.log_std[...])
class GaussianPolicy(d2l.ActorCritic):
"""The same interface over a Normal instead of a softmax; nothing that
consumes the interface changes."""
def __init__(self, obs_dim, act_dim, hidden=64, lr=1e-2, rngs=None):
rngs = nnx.Rngs(d2l.get_key()) if rngs is None else rngs
super().__init__(GaussianHead(obs_dim, act_dim, hidden, rngs),
nnx.Sequential(nnx.Linear(obs_dim, hidden, rngs=rngs),
jnp.tanh,
nnx.Linear(hidden, 1, rngs=rngs)), lr)
def log_prob(self, obs, act, policy=None):
mean, std = (self.policy if policy is None else policy)(obs)
return jax.scipy.stats.norm.logpdf(act, mean, std).sum(-1)
def act(self, obs, rng):
if not hasattr(self, '_fwd'): # compile the fixed-shape acting
self._fwd = nnx.cached_partial(nnx.jit(lambda net, o: net(o)),
self.policy) # forward, once
mean, std = self._fwd(jnp.asarray(obs))
return np.asarray(mean) + np.asarray(std) * rng.standard_normal(
mean.shape, dtype=np.float32)
def act_greedy(self, obs, rng=None):
return np.asarray(self.policy(jnp.asarray(obs))[0])