This commit is contained in:
wty-yy
2025-12-06 17:11:40 +08:00
parent a05d047b39
commit 9d509829a5
19 changed files with 364 additions and 60 deletions

View File

@@ -38,6 +38,21 @@ def class_to_dict(obj) -> dict:
result[key] = element
return result
def set_seed(seed: int):
import os
import torch
import random
import numpy as np
assert seed >= 0, "Seed must be non-negative."
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def str2bool(v):
if v.lower() in ('yes', 'true', 't', 'y', '1'):
return True
@@ -48,12 +63,21 @@ def str2bool(v):
def parse_args():
parser = ArgumentParser()
parameters = [
# Single run parameters
{"name": "--task-name", "type": str, "default": "base", "help": "Name of the task to run."},
{"name": "--experiment-name", "type": str, "help": "Name of the experiment to run."},
{"name": "--run-name", "type": str, "default": "run1", "help": "Name of the run."},
{"name": "--run-name", "type": str, "default": "run", "help": "Name of the run."},
{"name": "--model-path", "type": str, "help": "Path to the model file."},
{"name": "--headless", "action": "store_true", "default": False, "help": "Run in headless mode."},
{"name": "--save-video", "action": "store_true", "default": False, "help": "Save video output."},
{"name": "--seed", "type": int, "default": 42, "help": "Random seed."},
# Multiprocessing parameters, with different seeds
{"name": "--multi", "action": "store_true", "default": False, "help": "Enable multiprocessing."},
{"name": "--num-processes", "type": int, "default": 2, "help": "Number of parallel processes."},
{"name": "--seeds", "type": int, "nargs": "+", "default": [0], "help": "List of random seeds for multiple runs."},
{"name": "--base-masses", "type": float, "nargs": "+", "default": [-1, 0, 1], "help": "List of base masses for the model."},
{"name": "--frictions", "type": float, "nargs": "+", "default": [0.5, 1.0, 1.5], "help": "List of friction coefficients for the model."}
]
for param in parameters:
parser.add_argument(param['name'], **{k: v for k, v in param.items() if k != 'name'})

View File

@@ -7,7 +7,9 @@
@Blog : https://wty-yy.github.io/
@Desc : Task Registration Utility
'''
from robogauge.utils.helpers import parse_args
from robogauge import ROBOGAUGE_ROOT_DIR
from robogauge.utils.logger import logger
from robogauge.utils.helpers import parse_args, set_seed
class TaskRegister():
def __init__(self):
@@ -29,13 +31,13 @@ class TaskRegister():
def get_cfgs(self, name):
if name not in self.sim_cfgs:
raise ValueError(f"Task '{name}' is not registered.")
raise ValueError(f"Task '{name}' is not registered, checkout '{ROBOGAUGE_ROOT_DIR}/robogauge/tasks/__init__.py'.")
sim_cfg = self.sim_cfgs[name]
gauger_cfg = self.gauger_cfgs[name]
robot_cfg = self.robot_cfgs[name]
return sim_cfg, gauger_cfg, robot_cfg
def make_pipeline(self, args=None, sim_cfg=None, gauger_cfg=None, robot_cfg=None):
def make_pipeline(self, args=None, sim_cfg=None, gauger_cfg=None, robot_cfg=None, create_logger=True):
if args is None:
args = parse_args()
default_cfgs = self.get_cfgs(args.task_name)
@@ -48,7 +50,11 @@ class TaskRegister():
if args is not None:
self.update_args_to_cfg(sim_cfg, gauger_cfg, robot_cfg, args)
pipeline_class = self.get_pipeline_class(args.task_name)
return pipeline_class(args.run_name, sim_cfg, robot_cfg, gauger_cfg)
set_seed(args.seed)
run_name = args.run_name + f'_{args.seed}'
if create_logger:
logger.create(args.experiment_name, run_name)
return pipeline_class(run_name, sim_cfg, robot_cfg, gauger_cfg)
def update_args_to_cfg(self, sim_cfg, gauger_cfg, robot_cfg, args):
if args.model_path is not None: