MCPcopy Create free account
hub / github.com/LeCAR-Lab/model-based-diffusion / reverse_once

Function reverse_once

mbd/planners/mbd_planner.py:98–135  ·  view source on GitHub ↗
(carry, unused)

Source from the content-addressed store, hash-verified

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):

Callers 1

reverseFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected