Repository navigation
Expand file tree
/
Copy pathnemo_media.py
More file actions
92 lines (76 loc) · 2.85 KB
/
Copy pathnemo_media.py
File metadata and controls
92 lines (76 loc) · 2.85 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
import jax
from brax import envs
from brax.io import model
import mujoco
import jax.numpy as jnp
import mediapy
OBS_SIZE = 334
ACT_SIZE = 24
DT = 0.01
def generate_rollout(lstm=True):
# Import check
if lstm:
from nemo_lstm import NemoEnv
else:
from nemo_env_pd import NemoEnv
# Loading xml models
model_n = mujoco.MjModel.from_xml_path("nemo4b/scene.xml")
pelvis_b_id = mujoco.mj_name2id(model_n, mujoco.mjtObj.mjOBJ_SITE, 'pelvis_back')
pelvis_f_id = mujoco.mj_name2id(model_n, mujoco.mjtObj.mjOBJ_SITE, 'pelvis_front')
envs.register_environment('nemo', NemoEnv)
env = envs.create(env_name='nemo')
# JIT compile core functions
jit_reset = jax.jit(env.reset)
jit_step = jax.jit(env.step)
# Initialize state
state = jit_reset(jax.random.PRNGKey(0))
frames = [] # Store frames for video
# Load policy
saved_params = model.load_params('policies/walk_policy_acc7')
rng = jax.random.PRNGKey(0)
# Setup inference function
def makeIFN():
from brax.training.agents.ppo import networks as ppo_networks
from networks.lstm import make_ppo_networks
import functools
from brax.training.acme import running_statistics
mpn = make_ppo_networks
network_factory = functools.partial(
mpn,
policy_hidden_layer_sizes=(512, 256, 256, 128))
# normalize = running_statistics.normalize
normalize = lambda x, y: x
obs_size = OBS_SIZE
ppo_network = network_factory(
obs_size, ACT_SIZE, preprocess_observations_fn=normalize
)
make_inference_fn = ppo_networks.make_inference_fn(ppo_network)
return make_inference_fn
inference_fn = makeIFN()(saved_params)
jit_inference_fn = jax.jit(inference_fn)
# Run simulation
n_steps = 20000
for i in range(n_steps):
# Update state info
state.info["velocity"] = jax.numpy.array([0.4, 0.0])
# Calculate facing direction
data = state.pipeline_state
pp1 = data.site_xpos[pelvis_f_id]
pp2 = data.site_xpos[pelvis_b_id]
facing_vec = (pp1 - pp2)[0:2]
facing_vec = facing_vec / jnp.linalg.norm(facing_vec)
state.info["angvel"] = facing_vec[1] * -2
# Get action and step environment
act_rng, rng = jax.random.split(rng)
action, _ = jit_inference_fn(state.obs, act_rng)
state = jit_step(state, action)
# Store frame
frames.append(env.render(state.pipeline_state))
# Save video using mediapy
mediapy.write_video('nemo_simulation.mp4', frames, fps=60)
# Could also do show_video
# mediapy.show_video(frames, camera='track'), fps=1.0 / env.dt / render_every)
return frames
if __name__ == "__main__":
frames = generate_rollout(lstm=True)
print("Simulation complete! Video saved as 'nemo_simulation.mp4'")