| 96 | |
| 97 | @jax.jit |
| 98 | def reverse_once(carry, unused): |
| 99 | i, rng, Ybar_i = carry |
| 100 | Yi = Ybar_i * jnp.sqrt(alphas_bar[i]) |
| 101 | |
| 102 | # sample from q_i |
| 103 | rng, Y0s_rng = jax.random.split(rng) |
| 104 | eps_u = jax.random.normal(Y0s_rng, (args.Nsample, args.Hsample, Nu)) |
| 105 | Y0s = eps_u * sigmas[i] + Ybar_i |
| 106 | Y0s = jnp.clip(Y0s, -1.0, 1.0) |
| 107 | |
| 108 | # esitimate mu_0tm1 |
| 109 | rewss, qs = jax.vmap(rollout_us, in_axes=(None, 0))(state_init, Y0s) |
| 110 | rews = rewss.mean(axis=-1) |
| 111 | rew_std = rews.std() |
| 112 | rew_std = jnp.where(rew_std < 1e-4, 1.0, rew_std) |
| 113 | rew_mean = rews.mean() |
| 114 | logp0 = (rews - rew_mean) / rew_std / args.temp_sample |
| 115 | |
| 116 | # evalulate demo |
| 117 | if args.enable_demo: |
| 118 | xref_logpds = jax.vmap(env.eval_xref_logpd)(qs) |
| 119 | xref_logpds = xref_logpds - xref_logpds.max() |
| 120 | logpdemo = ( |
| 121 | (xref_logpds + env.rew_xref - rew_mean) / rew_std / args.temp_sample |
| 122 | ) |
| 123 | demo_mask = logpdemo > logp0 |
| 124 | logp0 = jnp.where(demo_mask, logpdemo, logp0) |
| 125 | logp0 = (logp0 - logp0.mean()) / logp0.std() / args.temp_sample |
| 126 | |
| 127 | weights = jax.nn.softmax(logp0) |
| 128 | Ybar = jnp.einsum("n,nij->ij", weights, Y0s) # NOTE: update only with reward |
| 129 | |
| 130 | score = 1 / (1.0 - alphas_bar[i]) * (-Yi + jnp.sqrt(alphas_bar[i]) * Ybar) |
| 131 | Yim1 = 1 / jnp.sqrt(alphas[i]) * (Yi + (1.0 - alphas_bar[i]) * score) |
| 132 | |
| 133 | Ybar_im1 = Yim1 / jnp.sqrt(alphas_bar[i - 1]) |
| 134 | |
| 135 | return (i - 1, rng, Ybar_im1), rews.mean() |
| 136 | |
| 137 | # run reverse |
| 138 | def reverse(YN, rng): |