(self, q)
| 87 | |
| 88 | @partial(jax.jit, static_argnums=(0,)) |
| 89 | def get_reward(self, q): |
| 90 | reward = ( |
| 91 | 1.0 - (jnp.clip(jnp.linalg.norm(q[:2] - self.xg[:2]), 0.0, 0.2) / 0.2) ** 2 |
| 92 | ) |
| 93 | return reward |
| 94 | |
| 95 | @partial(jax.jit, static_argnums=(0,)) |
| 96 | def eval_xref_logpd(self, xs): |