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

Function run_diffusion

mbd/planners/mbd_planner.py:38–182  ·  view source on GitHub ↗
(args: Args)

Source from the content-addressed store, hash-verified

36
37
38def run_diffusion(args: Args):
39
40 rng = jax.random.PRNGKey(seed=args.seed)
41
42 ## setup env
43
44 # recommended temperature for envs
45 temp_recommend = {
46 "ant": 0.1,
47 "halfcheetah": 0.4,
48 "hopper": 0.1,
49 "humanoidstandup": 0.1,
50 "humanoidrun": 0.1,
51 "walker2d": 0.1,
52 "pushT": 0.2,
53 }
54 Ndiffuse_recommend = {
55 "pushT": 200,
56 "humanoidrun": 300,
57 }
58 Nsample_recommend = {
59 "humanoidrun": 8192,
60 }
61 Hsample_recommend = {
62 "pushT": 40,
63 }
64 if not args.disable_recommended_params:
65 args.temp_sample = temp_recommend.get(args.env_name, args.temp_sample)
66 args.Ndiffuse = Ndiffuse_recommend.get(args.env_name, args.Ndiffuse)
67 args.Nsample = Nsample_recommend.get(args.env_name, args.Nsample)
68 args.Hsample = Hsample_recommend.get(args.env_name, args.Hsample)
69 print(f"override temp_sample to {args.temp_sample}")
70 env = mbd.envs.get_env(args.env_name)
71 Nx = env.observation_size
72 Nu = env.action_size
73 # env functions
74 step_env_jit = jax.jit(env.step)
75 reset_env_jit = jax.jit(env.reset)
76 # eval_us = jax.jit(functools.partial(mbd.utils.eval_us, step_env_jit))
77 rollout_us = jax.jit(functools.partial(mbd.utils.rollout_us, step_env_jit))
78
79 rng, rng_reset = jax.random.split(rng) # NOTE: rng_reset should never be changed.
80 state_init = reset_env_jit(rng_reset)
81
82 ## run diffusion
83
84 betas = jnp.linspace(args.beta0, args.betaT, args.Ndiffuse)
85 alphas = 1.0 - betas
86 alphas_bar = jnp.cumprod(alphas)
87 sigmas = jnp.sqrt(1 - alphas_bar)
88 Sigmas_cond = (
89 (1 - alphas) * (1 - jnp.sqrt(jnp.roll(alphas_bar, 1))) / (1 - alphas_bar)
90 )
91 sigmas_cond = jnp.sqrt(Sigmas_cond)
92 sigmas_cond = sigmas_cond.at[0].set(0.0)
93 print(f"init sigma = {sigmas[-1]:.2e}")
94
95 YN = jnp.zeros([args.Hsample, Nu])

Callers 1

mbd_planner.pyFile · 0.85

Calls 4

reverseFunction · 0.85
rollout_usFunction · 0.85
renderMethod · 0.80
render_usFunction · 0.50

Tested by

no test coverage detected