v0.1.9; wait robogauge eval after all training; ppo default train 150k, add yaml cfg save
This commit is contained in:
@@ -1,3 +1,7 @@
|
|||||||
|
# 20260112
|
||||||
|
## v0.1.9
|
||||||
|
1. 添加训练完成时, 等待RoboGauge全部评估完全后结束
|
||||||
|
2. PPO训练步数对齐到150k, 并加入配置yaml自动存储功能
|
||||||
# 20260109
|
# 20260109
|
||||||
## v0.1.8
|
## v0.1.8
|
||||||
1. 加入两个配置legged_gym/envs/go2/go2_config_vanilla.py和legged_gym/envs/go2/go2_config_vanilla2.py
|
1. 加入两个配置legged_gym/envs/go2/go2_config_vanilla.py和legged_gym/envs/go2/go2_config_vanilla2.py
|
||||||
|
|||||||
@@ -393,6 +393,9 @@ class LeggedRobotCfgREMCTS(LeggedRobotCfgCTS):
|
|||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
student_encoder_hidden_dims = [512, 256, 128]
|
student_encoder_hidden_dims = [512, 256, 128]
|
||||||
|
|
||||||
|
class algorithm(LeggedRobotCfgCTS.algorithm):
|
||||||
|
load_balance_coef = 0.01 # coefficient for load balance loss
|
||||||
|
|
||||||
class runner(LeggedRobotCfgCTS.runner):
|
class runner(LeggedRobotCfgCTS.runner):
|
||||||
policy_class_name = 'ActorCriticREMCTS'
|
policy_class_name = 'ActorCriticREMCTS'
|
||||||
algorithm_class_name = 'REMCTS'
|
algorithm_class_name = 'REMCTS'
|
||||||
@@ -238,7 +238,7 @@ class GO2CfgPPO(LeggedRobotCfgPPO):
|
|||||||
class runner(LeggedRobotCfgPPO.runner):
|
class runner(LeggedRobotCfgPPO.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_ppo'
|
experiment_name = 'go2_ppo'
|
||||||
max_iterations = 100000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgCTS(LeggedRobotCfgCTS):
|
class GO2CfgCTS(LeggedRobotCfgCTS):
|
||||||
|
|||||||
@@ -268,7 +268,7 @@ class GO2CfgPPO(LeggedRobotCfgPPO):
|
|||||||
class runner(LeggedRobotCfgPPO.runner):
|
class runner(LeggedRobotCfgPPO.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_ppo'
|
experiment_name = 'go2_ppo'
|
||||||
max_iterations = 100000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgCTS(LeggedRobotCfgCTS):
|
class GO2CfgCTS(LeggedRobotCfgCTS):
|
||||||
|
|||||||
@@ -253,7 +253,7 @@ class GO2CfgPPO(LeggedRobotCfgPPO):
|
|||||||
class runner(LeggedRobotCfgPPO.runner):
|
class runner(LeggedRobotCfgPPO.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_ppo'
|
experiment_name = 'go2_ppo'
|
||||||
max_iterations = 100000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgCTS(LeggedRobotCfgCTS):
|
class GO2CfgCTS(LeggedRobotCfgCTS):
|
||||||
|
|||||||
@@ -253,7 +253,7 @@ class GO2CfgPPO(LeggedRobotCfgPPO):
|
|||||||
class runner(LeggedRobotCfgPPO.runner):
|
class runner(LeggedRobotCfgPPO.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_ppo'
|
experiment_name = 'go2_ppo'
|
||||||
max_iterations = 100000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgCTS(LeggedRobotCfgCTS):
|
class GO2CfgCTS(LeggedRobotCfgCTS):
|
||||||
|
|||||||
@@ -33,6 +33,8 @@ import os
|
|||||||
from collections import deque
|
from collections import deque
|
||||||
import statistics
|
import statistics
|
||||||
import yaml
|
import yaml
|
||||||
|
import numpy as np
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from torch.utils.tensorboard import SummaryWriter
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
import torch
|
import torch
|
||||||
@@ -40,8 +42,20 @@ import torch
|
|||||||
from rsl_rl.algorithms import PPO
|
from rsl_rl.algorithms import PPO
|
||||||
from rsl_rl.modules import ActorCritic, ActorCriticRecurrent
|
from rsl_rl.modules import ActorCritic, ActorCriticRecurrent
|
||||||
from rsl_rl.env import VecEnv
|
from rsl_rl.env import VecEnv
|
||||||
|
from legged_gym.utils.helpers import class_to_dict
|
||||||
from legged_gym.utils.exporter import export_policy_as_jit
|
from legged_gym.utils.exporter import export_policy_as_jit
|
||||||
|
|
||||||
|
def numpy_representer(dumper, data):
|
||||||
|
return dumper.represent_float(float(data))
|
||||||
|
|
||||||
|
def numpy_int_representer(dumper, data):
|
||||||
|
return dumper.represent_int(int(data))
|
||||||
|
|
||||||
|
# Add the numpy representer to yaml
|
||||||
|
yaml.add_representer(np.float32, numpy_representer, Dumper=yaml.SafeDumper)
|
||||||
|
yaml.add_representer(np.float64, numpy_representer, Dumper=yaml.SafeDumper)
|
||||||
|
yaml.add_representer(np.int32, numpy_int_representer, Dumper=yaml.SafeDumper)
|
||||||
|
yaml.add_representer(np.int64, numpy_int_representer, Dumper=yaml.SafeDumper)
|
||||||
|
|
||||||
class OnPolicyRunner:
|
class OnPolicyRunner:
|
||||||
|
|
||||||
@@ -81,6 +95,10 @@ class OnPolicyRunner:
|
|||||||
self.current_learning_iteration = 0
|
self.current_learning_iteration = 0
|
||||||
|
|
||||||
_, _ = self.env.reset()
|
_, _ = self.env.reset()
|
||||||
|
if self.log_dir is not None and self.env.cfg.env.test is False:
|
||||||
|
Path(self.log_dir).mkdir(parents=True, exist_ok=True)
|
||||||
|
all_cfg = {"train_cfg": train_cfg, "env_cfg": class_to_dict(self.env.cfg)}
|
||||||
|
yaml.safe_dump(all_cfg, open(os.path.join(self.log_dir, 'config.yaml'), 'w'))
|
||||||
|
|
||||||
# robogauge client
|
# robogauge client
|
||||||
try:
|
try:
|
||||||
@@ -144,11 +162,11 @@ class OnPolicyRunner:
|
|||||||
if self.log_dir is not None:
|
if self.log_dir is not None:
|
||||||
self.log(locals())
|
self.log(locals())
|
||||||
if it % self.save_interval == 0:
|
if it % self.save_interval == 0:
|
||||||
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(it)), it)
|
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(it)), it, False)
|
||||||
ep_infos.clear()
|
ep_infos.clear()
|
||||||
|
|
||||||
self.current_learning_iteration += num_learning_iterations
|
self.current_learning_iteration += num_learning_iterations
|
||||||
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(self.current_learning_iteration)), it)
|
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(self.current_learning_iteration)), it, True)
|
||||||
|
|
||||||
def log(self, locs, width=80, pad=35):
|
def log(self, locs, width=80, pad=35):
|
||||||
self.tot_timesteps += self.num_steps_per_env * self.env.num_envs
|
self.tot_timesteps += self.num_steps_per_env * self.env.num_envs
|
||||||
@@ -219,20 +237,20 @@ class OnPolicyRunner:
|
|||||||
locs['num_learning_iterations'] - locs['it']):.1f}s\n""")
|
locs['num_learning_iterations'] - locs['it']):.1f}s\n""")
|
||||||
print(log_string)
|
print(log_string)
|
||||||
|
|
||||||
def save(self, path, it, infos=None):
|
def save(self, path, it, last_model, infos=None):
|
||||||
torch.save({
|
torch.save({
|
||||||
'model_state_dict': self.alg.actor_critic.state_dict(),
|
'model_state_dict': self.alg.actor_critic.state_dict(),
|
||||||
'optimizer_state_dict': self.alg.optimizer.state_dict(),
|
'optimizer_state_dict': self.alg.optimizer.state_dict(),
|
||||||
'iter': self.current_learning_iteration,
|
'iter': self.current_learning_iteration,
|
||||||
'infos': infos,
|
'infos': infos,
|
||||||
}, path)
|
}, path)
|
||||||
self.update_robogauge(it)
|
self.update_robogauge(it, last_model)
|
||||||
|
|
||||||
def update_robogauge(self, it):
|
def update_robogauge(self, it, last_model):
|
||||||
if self.robogauge_client is None:
|
if self.robogauge_client is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
if it % 500 == 0:
|
if it % 500 == 0 or last_model:
|
||||||
# export jit model
|
# export jit model
|
||||||
jit_dir = os.path.join(self.log_dir, 'jit_models')
|
jit_dir = os.path.join(self.log_dir, 'jit_models')
|
||||||
jit_path = os.path.join(jit_dir, f'policy_jit_{it}.pt')
|
jit_path = os.path.join(jit_dir, f'policy_jit_{it}.pt')
|
||||||
@@ -245,17 +263,33 @@ class OnPolicyRunner:
|
|||||||
task_name=task_name,
|
task_name=task_name,
|
||||||
experiment_name=self.cfg["experiment_name"]
|
experiment_name=self.cfg["experiment_name"]
|
||||||
)
|
)
|
||||||
self.robogauge_client.monitor_tasks()
|
check_times = 1
|
||||||
results_dir = os.path.join(self.log_dir, 'robogauge_results')
|
if last_model:
|
||||||
os.makedirs(results_dir, exist_ok=True)
|
check_times = int(1e9) # keep checking until the last model is evaluated
|
||||||
for task_id, resp in self.robogauge_client.response_data.items():
|
while check_times > 0:
|
||||||
scores = resp['results']['scores']
|
check_times -= 1
|
||||||
step = resp['step']
|
self.robogauge_client.monitor_tasks()
|
||||||
for key, val in scores.items():
|
results_dir = os.path.join(self.log_dir, 'robogauge_results')
|
||||||
self.writer.add_scalar(f'RoboGauge/{key}', val, step)
|
os.makedirs(results_dir, exist_ok=True)
|
||||||
results_path = os.path.join(results_dir, f'results_{step}.yaml')
|
result_received = False
|
||||||
with open(results_path, 'w', encoding='utf-8') as f:
|
for task_id, resp in self.robogauge_client.response_data.items():
|
||||||
yaml.dump(resp['results'], f, allow_unicode=True, sort_keys=False)
|
scores = resp['results']['scores']
|
||||||
|
step = resp['step']
|
||||||
|
if step == it:
|
||||||
|
result_received = True
|
||||||
|
for key, val in scores.items():
|
||||||
|
self.writer.add_scalar(f'RoboGauge/{key}', val, step)
|
||||||
|
results_path = os.path.join(results_dir, f'results_{step}.yaml')
|
||||||
|
with open(results_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(resp['results'], f, allow_unicode=True, sort_keys=False)
|
||||||
|
|
||||||
|
if last_model and result_received:
|
||||||
|
print(f"RoboGauge result for step {it} received. Exiting wait loop.")
|
||||||
|
break
|
||||||
|
|
||||||
|
if check_times > 0:
|
||||||
|
print("Sleeping for 1 minute before checking RoboGauge results again...")
|
||||||
|
time.sleep(60) # wait for 1 minute before checking again
|
||||||
|
|
||||||
def load(self, path, load_optimizer=True):
|
def load(self, path, load_optimizer=True):
|
||||||
loaded_dict = torch.load(path)
|
loaded_dict = torch.load(path)
|
||||||
|
|||||||
@@ -193,10 +193,10 @@ class OnPolicyRunnerCTS:
|
|||||||
if self.log_dir is not None:
|
if self.log_dir is not None:
|
||||||
self.log(locals())
|
self.log(locals())
|
||||||
if it % self.save_interval == 0:
|
if it % self.save_interval == 0:
|
||||||
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(it)), it)
|
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(it)), it, False)
|
||||||
ep_infos.clear()
|
ep_infos.clear()
|
||||||
|
|
||||||
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(self.current_learning_iteration)), it)
|
self.save(os.path.join(self.log_dir, 'model_{}.pt'.format(self.current_learning_iteration)), it, True)
|
||||||
|
|
||||||
def log(self, locs, width=80, pad=35):
|
def log(self, locs, width=80, pad=35):
|
||||||
self.tot_timesteps += self.num_steps_per_env * self.env.num_envs
|
self.tot_timesteps += self.num_steps_per_env * self.env.num_envs
|
||||||
@@ -281,7 +281,7 @@ class OnPolicyRunnerCTS:
|
|||||||
locs['tot_iter'] - locs['it']):.1f}s\n""")
|
locs['tot_iter'] - locs['it']):.1f}s\n""")
|
||||||
print(log_string)
|
print(log_string)
|
||||||
|
|
||||||
def save(self, path, it, infos=None):
|
def save(self, path, it, last_model, infos=None):
|
||||||
torch.save({
|
torch.save({
|
||||||
'model_state_dict': self.alg.model.state_dict(),
|
'model_state_dict': self.alg.model.state_dict(),
|
||||||
'optimizer1_state_dict': self.alg.optimizer1.state_dict(),
|
'optimizer1_state_dict': self.alg.optimizer1.state_dict(),
|
||||||
@@ -289,13 +289,13 @@ class OnPolicyRunnerCTS:
|
|||||||
'iter': self.current_learning_iteration,
|
'iter': self.current_learning_iteration,
|
||||||
'infos': infos,
|
'infos': infos,
|
||||||
}, path)
|
}, path)
|
||||||
self.update_robogauge(it)
|
self.update_robogauge(it, last_model)
|
||||||
|
|
||||||
def update_robogauge(self, it):
|
def update_robogauge(self, it, last_model):
|
||||||
if self.robogauge_client is None:
|
if self.robogauge_client is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
if it % 500 == 0:
|
if it % 500 == 0 or last_model:
|
||||||
# export jit model
|
# export jit model
|
||||||
jit_dir = os.path.join(self.log_dir, 'jit_models')
|
jit_dir = os.path.join(self.log_dir, 'jit_models')
|
||||||
jit_path = os.path.join(jit_dir, f'policy_jit_{it}.pt')
|
jit_path = os.path.join(jit_dir, f'policy_jit_{it}.pt')
|
||||||
@@ -308,17 +308,33 @@ class OnPolicyRunnerCTS:
|
|||||||
task_name=task_name,
|
task_name=task_name,
|
||||||
experiment_name=self.cfg["experiment_name"]
|
experiment_name=self.cfg["experiment_name"]
|
||||||
)
|
)
|
||||||
self.robogauge_client.monitor_tasks()
|
check_times = 1
|
||||||
results_dir = os.path.join(self.log_dir, 'robogauge_results')
|
if last_model:
|
||||||
os.makedirs(results_dir, exist_ok=True)
|
check_times = int(1e9) # keep checking until manually stopped
|
||||||
for task_id, resp in self.robogauge_client.response_data.items():
|
while check_times > 0:
|
||||||
scores = resp['results']['scores']
|
check_times -= 1
|
||||||
step = resp['step']
|
self.robogauge_client.monitor_tasks()
|
||||||
for key, val in scores.items():
|
results_dir = os.path.join(self.log_dir, 'robogauge_results')
|
||||||
self.writer.add_scalar(f'RoboGauge/{key}', val, step)
|
os.makedirs(results_dir, exist_ok=True)
|
||||||
results_path = os.path.join(results_dir, f'results_{step}.yaml')
|
result_received = False
|
||||||
with open(results_path, 'w', encoding='utf-8') as f:
|
for task_id, resp in self.robogauge_client.response_data.items():
|
||||||
yaml.dump(resp['results'], f, allow_unicode=True, sort_keys=False)
|
scores = resp['results']['scores']
|
||||||
|
step = resp['step']
|
||||||
|
if step == it:
|
||||||
|
result_received = True
|
||||||
|
for key, val in scores.items():
|
||||||
|
self.writer.add_scalar(f'RoboGauge/{key}', val, step)
|
||||||
|
results_path = os.path.join(results_dir, f'results_{step}.yaml')
|
||||||
|
with open(results_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(resp['results'], f, allow_unicode=True, sort_keys=False)
|
||||||
|
|
||||||
|
if last_model and result_received:
|
||||||
|
print(f"RoboGauge result for step {it} received. Exiting wait loop.")
|
||||||
|
break
|
||||||
|
|
||||||
|
if check_times > 0:
|
||||||
|
print("Sleeping for 1 minute before checking RoboGauge results again...")
|
||||||
|
time.sleep(60) # wait for 1 minute before checking again
|
||||||
|
|
||||||
def load(self, path, load_optimizer=True):
|
def load(self, path, load_optimizer=True):
|
||||||
loaded_dict = torch.load(path)
|
loaded_dict = torch.load(path)
|
||||||
|
|||||||
2
setup.py
2
setup.py
@@ -2,7 +2,7 @@ from setuptools import find_packages
|
|||||||
from distutils.core import setup
|
from distutils.core import setup
|
||||||
|
|
||||||
setup(name='go2_rl_gym',
|
setup(name='go2_rl_gym',
|
||||||
version='0.1.8',
|
version='0.1.9',
|
||||||
author='Wu Tianyang',
|
author='Wu Tianyang',
|
||||||
license="MIT",
|
license="MIT",
|
||||||
packages=find_packages(),
|
packages=find_packages(),
|
||||||
|
|||||||
Reference in New Issue
Block a user