mirror of
https://github.com/wassname/kair_algorithms_draft.git
synced 2026-10-04 12:40:46 +08:00
* Convert code format to python2.7 (SAC) * Convert code format python2.7 (TD3, all fD) * Remove no use import and black setting * Change SAC param * Change env name Reacher-v2 to v1 * Remove old version reacher training script * Convert code format python2.7 * Modify .travis.yml * Add install command python3.6 & black on Makefile * Fix seperator to tab on Makefile * Modify Makefile * Fix little error * Change td3 gamma parameter
71 lines
2.0 KiB
Python
71 lines
2.0 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Train or test algorithms on Reacher-v1 of Mujoco.
|
|
|
|
- Author: Kh Kim
|
|
- Contact: kh.kim@medipixel.io
|
|
"""
|
|
|
|
import argparse
|
|
import importlib
|
|
|
|
import gym
|
|
|
|
import algorithms.common.helper_functions as common_utils
|
|
|
|
# configurations
|
|
parser = argparse.ArgumentParser(description="Pytorch RL algorithms")
|
|
parser.add_argument(
|
|
"--seed", type=int, default=777, help="random seed for reproducibility"
|
|
)
|
|
parser.add_argument("--algo", type=str, default="sac", help="choose an algorithm")
|
|
parser.add_argument(
|
|
"--test", dest="test", action="store_true", help="test mode (no training)"
|
|
)
|
|
parser.add_argument(
|
|
"--load-from", type=str, help="load the saved model and optimizer at the beginning"
|
|
)
|
|
parser.add_argument(
|
|
"--off-render", dest="render", action="store_false", help="turn off rendering"
|
|
)
|
|
parser.add_argument(
|
|
"--render-after",
|
|
type=int,
|
|
default=0,
|
|
help="start rendering after the input number of episode",
|
|
)
|
|
parser.add_argument("--log", dest="log", action="store_true", help="turn on logging")
|
|
parser.add_argument("--save-period", type=int, default=200, help="save model period")
|
|
parser.add_argument("--episode-num", type=int, default=20000, help="total episode num")
|
|
parser.add_argument(
|
|
"--max-episode-steps", type=int, default=-1, help="max episode step"
|
|
)
|
|
parser.add_argument(
|
|
"--demo-path", type=str, default="data/reacher_demo.pkl", help="demonstration path"
|
|
)
|
|
|
|
parser.set_defaults(test=False)
|
|
parser.set_defaults(load_from=None)
|
|
parser.set_defaults(render=True)
|
|
parser.set_defaults(log=False)
|
|
args = parser.parse_args()
|
|
|
|
|
|
def main():
|
|
"""Main."""
|
|
# env initialization
|
|
env = gym.make("Reacher-v1")
|
|
state_dim = env.observation_space.shape[0]
|
|
action_dim = env.action_space.shape[0]
|
|
|
|
# set a random seed
|
|
common_utils.set_random_seed(args.seed, env)
|
|
|
|
# run
|
|
module_path = "examples.reacher-v1." + args.algo
|
|
example = importlib.import_module(module_path)
|
|
example.run(env, args, state_dim, action_dim)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|