pytorch version

This commit is contained in:
Kolesnikov Sergey
2017-11-15 22:18:46 +03:00
parent 34993abdf7
commit 7401266fe7
49 changed files with 5435 additions and 1 deletions
View File
+70
View File
@@ -0,0 +1,70 @@
import os
import torch
import copy
from multiprocessing import Value
from common.misc_util import str2params, create_if_need
from common.env_wrappers import create_env
from common.torch_util import activations, hard_update
from ddpg.model import create_model, create_act_update_fns, train_multi_thread
from ddpg.train import parse_args
def debug(args, model_fn, act_update_fns, multi_thread):
create_if_need(args.logdir)
env = create_env(args)
if args.flip_state_action and hasattr(env, "state_transform"):
args.flip_states = env.state_transform.flip_states
args.n_action = env.action_space.shape[0]
args.n_observation = env.observation_space.shape[0]
args.actor_layers = str2params(args.actor_layers)
args.critic_layers = str2params(args.critic_layers)
args.actor_activation = activations[args.actor_activation]
args.critic_activation = activations[args.critic_activation]
actor, critic = model_fn(args)
if args.restore_actor_from is not None:
actor.load_state_dict(torch.load(args.restore_actor_from))
if args.restore_critic_from is not None:
critic.load_state_dict(torch.load(args.restore_critic_from))
actor.train()
critic.train()
actor.share_memory()
critic.share_memory()
target_actor = copy.deepcopy(actor)
target_critic = copy.deepcopy(critic)
hard_update(target_actor, actor)
hard_update(target_critic, critic)
target_actor.train()
critic.train()
target_actor.share_memory()
target_critic.share_memory()
_, _, save_fn = act_update_fns(actor, critic, target_actor, target_critic, args)
args.thread = 0
best_reward = Value("f", 0.0)
multi_thread(actor, critic, target_actor, target_critic, args, act_update_fns, best_reward)
save_fn()
if __name__ == '__main__':
os.environ['OMP_NUM_THREADS'] = '1'
torch.set_num_threads(1)
args = parse_args()
debug(
args,
create_model,
create_act_update_fns,
train_multi_thread)
+477
View File
@@ -0,0 +1,477 @@
import random
import numpy as np
import torch
import queue as py_queue
import time
import torch.nn as nn
from pprint import pprint
from ddpg.nets import Actor, Critic
from common.torch_util import to_numpy, to_tensor, soft_update
from common.misc_util import create_if_need, set_global_seeds
from common.logger import Logger
from common.buffers import create_buffer
from common.loss import create_loss, create_decay_fn
from common.env_wrappers import create_env
from common.random_process import create_random_process
def create_model(args):
actor = Actor(
args.n_observation, args.n_action, args.actor_layers,
activation=args.actor_activation,
layer_norm=args.actor_layer_norm,
parameters_noise=args.actor_parameters_noise,
parameters_noise_factorised=args.actor_parameters_noise_factorised,
last_activation=nn.Tanh)
critic = Critic(
args.n_observation, args.n_action, args.critic_layers,
activation=args.critic_activation,
layer_norm=args.critic_layer_norm,
parameters_noise=args.critic_parameters_noise,
parameters_noise_factorised=args.critic_parameters_noise_factorised)
pprint(actor)
pprint(critic)
return actor, critic
def create_act_update_fns(actor, critic, target_actor, target_critic, args):
actor_optim = torch.optim.Adam(actor.parameters(), lr=args.actor_lr)
critic_optim = torch.optim.Adam(critic.parameters(), lr=args.critic_lr)
criterion = create_loss(args)
low_action_boundary = -1.
high_action_boundary = 1.
def act_fn(observation, noise=0):
nonlocal actor
action = to_numpy(actor(to_tensor(np.array([observation], dtype=np.float32)))).squeeze(0)
action += noise
action = np.clip(action, low_action_boundary, high_action_boundary)
return action
def update_fn(
observations, actions, rewards, next_observations, dones, weights,
actor_lr=1e-4, critic_lr=1e-3):
nonlocal actor, critic, target_actor, target_critic, actor_optim, critic_optim
if hasattr(args, "flip_states"):
observations_flip = args.flip_states(observations)
next_observations_flip = args.flip_states(next_observations)
actions_flip = np.zeros_like(actions)
actions_flip[:, :args.n_action // 2] = actions[:, args.n_action // 2:]
actions_flip[:, args.n_action // 2:] = actions[:, :args.n_action // 2]
observations = np.concatenate((observations, observations_flip))
actions = np.concatenate((actions, actions_flip))
rewards = np.tile(rewards.ravel(), 2)
next_observations = np.concatenate((next_observations, next_observations_flip))
dones = np.tile(dones.ravel(), 2)
dones = dones[:, None].astype(np.bool)
rewards = rewards[:, None].astype(np.float32)
dones = to_tensor(np.invert(dones).astype(np.float32))
rewards = to_tensor(rewards)
weights = to_tensor(weights, requires_grad=False)
next_v_values = target_critic(
to_tensor(next_observations, volatile=True),
target_actor(to_tensor(next_observations, volatile=True)),
)
next_v_values.volatile = False
reward_predicted = dones * args.gamma * next_v_values
td_target = rewards + reward_predicted
# Critic update
critic.zero_grad()
v_values = critic(to_tensor(observations), to_tensor(actions))
value_loss = criterion(v_values, td_target, weights=weights)
value_loss.backward()
torch.nn.utils.clip_grad_norm(critic.parameters(), args.grad_clip)
for param_group in critic_optim.param_groups:
param_group["lr"] = critic_lr
critic_optim.step()
# Actor update
actor.zero_grad()
policy_loss = -critic(
to_tensor(observations),
actor(to_tensor(observations))
)
policy_loss = torch.mean(policy_loss * weights)
policy_loss.backward()
torch.nn.utils.clip_grad_norm(actor.parameters(), args.grad_clip)
for param_group in actor_optim.param_groups:
param_group["lr"] = actor_lr
actor_optim.step()
# Target update
soft_update(target_actor, actor, args.tau)
soft_update(target_critic, critic, args.tau)
metrics = {
"value_loss": value_loss,
"policy_loss": policy_loss
}
td_v_values = critic(
to_tensor(observations, volatile=True, requires_grad=False),
to_tensor(actions, volatile=True, requires_grad=False))
td_error = td_target - td_v_values
info = {
"td_error": to_numpy(td_error)
}
return metrics, info
def save_fn(episode=None):
nonlocal actor, critic
if episode is None:
save_path = args.logdir
else:
save_path = "{}/episode_{}".format(args.logdir, episode)
create_if_need(save_path)
torch.save(actor.state_dict(), "{}/actor_state_dict.pkl".format(save_path))
torch.save(critic.state_dict(), "{}/critic_state_dict.pkl".format(save_path))
torch.save(target_actor.state_dict(), "{}/target_actor_state_dict.pkl".format(save_path))
torch.save(target_critic.state_dict(), "{}/target_critic_state_dict.pkl".format(save_path))
return act_fn, update_fn, save_fn
def train_multi_thread(actor, critic, target_actor, target_critic, args, prepare_fn, best_reward):
workerseed = args.seed + 241 * args.thread
set_global_seeds(workerseed)
args.logdir = "{}/thread_{}".format(args.logdir, args.thread)
create_if_need(args.logdir)
act_fn, update_fn, save_fn = prepare_fn(actor, critic, target_actor, target_critic, args)
logger = Logger(args.logdir)
buffer = create_buffer(args)
if args.prioritized_replay:
beta_deacy_fn = create_decay_fn(
"linear",
initial_value=args.prioritized_replay_beta0,
final_value=1.0,
max_step=args.max_episodes)
env = create_env(args)
random_process = create_random_process(args)
actor_learning_rate_decay_fn = create_decay_fn(
"linear",
initial_value=args.actor_lr,
final_value=args.actor_lr_end,
max_step=args.max_episodes)
critic_learning_rate_decay_fn = create_decay_fn(
"linear",
initial_value=args.critic_lr,
final_value=args.critic_lr_end,
max_step=args.max_episodes)
epsilon_cycle_len = random.randint(args.epsilon_cycle_len // 2, args.epsilon_cycle_len * 2)
epsilon_decay_fn = create_decay_fn(
"cycle",
initial_value=args.initial_epsilon,
final_value=args.final_epsilon,
cycle_len=epsilon_cycle_len,
num_cycles=args.max_episodes // epsilon_cycle_len)
episode = 0
step = 0
start_time = time.time()
while episode < args.max_episodes:
if episode % 100 == 0:
env = create_env(args)
seed = random.randrange(2 ** 32 - 2)
actor_lr = actor_learning_rate_decay_fn(episode)
critic_lr = critic_learning_rate_decay_fn(episode)
epsilon = min(args.initial_epsilon, max(args.final_epsilon, epsilon_decay_fn(episode)))
episode_metrics = {
"value_loss": 0.0,
"policy_loss": 0.0,
"reward": 0.0,
"step": 0,
"epsilon": epsilon
}
observation = env.reset(seed=seed, difficulty=args.difficulty)
random_process.reset_states()
done = False
while not done:
action = act_fn(observation, noise=epsilon*random_process.sample())
next_observation, reward, done, _ = env.step(action)
buffer.add(observation, action, reward, next_observation, done)
episode_metrics["reward"] += reward
episode_metrics["step"] += 1
if len(buffer) >= args.train_steps:
if args.prioritized_replay:
(tr_observations, tr_actions, tr_rewards, tr_next_observations, tr_dones,
weights, batch_idxes) = \
buffer.sample(batch_size=args.batch_size, beta=beta_deacy_fn(episode))
else:
(tr_observations, tr_actions, tr_rewards, tr_next_observations, tr_dones) = \
buffer.sample(batch_size=args.batch_size)
weights, batch_idxes = np.ones_like(tr_rewards), None
step_metrics, step_info = update_fn(
tr_observations, tr_actions, tr_rewards,
tr_next_observations, tr_dones,
weights, actor_lr, critic_lr)
if args.prioritized_replay:
new_priorities = np.abs(step_info["td_error"]) + 1e-6
buffer.update_priorities(batch_idxes, new_priorities)
for key, value in step_metrics.items():
value = to_numpy(value)[0]
episode_metrics[key] += value
observation = next_observation
episode += 1
if episode_metrics["reward"] > 15.0 * args.reward_scale \
and episode_metrics["reward"] > best_reward.value:
best_reward.value = episode_metrics["reward"]
logger.scalar_summary("best reward", best_reward.value, episode)
save_fn(episode)
step += episode_metrics["step"]
elapsed_time = time.time() - start_time
for key, value in episode_metrics.items():
value = value if "loss" not in key else value / episode_metrics["step"]
logger.scalar_summary(key, value, episode)
logger.scalar_summary(
"episode per minute",
episode / elapsed_time * 60,
episode)
logger.scalar_summary(
"step per second",
step / elapsed_time,
episode)
logger.scalar_summary("actor lr", actor_lr, episode)
logger.scalar_summary("critic lr", critic_lr, episode)
if episode % args.save_step == 0:
save_fn(episode)
if elapsed_time > 86400 * args.max_train_days:
episode = args.max_episodes + 1
save_fn(episode)
raise KeyboardInterrupt
def train_single_thread(
actor, critic, target_actor, target_critic, args, prepare_fn,
global_episode, global_update_step, episodes_queue):
workerseed = args.seed + 241 * args.thread
set_global_seeds(workerseed)
args.logdir = "{}/thread_{}".format(args.logdir, args.thread)
create_if_need(args.logdir)
_, update_fn, save_fn = prepare_fn(actor, critic, target_actor, target_critic, args)
logger = Logger(args.logdir)
buffer = create_buffer(args)
if args.prioritized_replay:
beta_deacy_fn = create_decay_fn(
"linear",
initial_value=args.prioritized_replay_beta0,
final_value=1.0,
max_step=args.max_update_steps)
actor_learning_rate_decay_fn = create_decay_fn(
"linear",
initial_value=args.actor_lr,
final_value=args.actor_lr_end,
max_step=args.max_update_steps)
critic_learning_rate_decay_fn = create_decay_fn(
"linear",
initial_value=args.critic_lr,
final_value=args.critic_lr_end,
max_step=args.max_update_steps)
update_step = 0
received_examples = 1 # just hack
while global_episode.value < args.max_episodes * (args.num_threads - args.num_train_threads) \
and global_update_step.value < args.max_update_steps * args.num_train_threads:
actor_lr = actor_learning_rate_decay_fn(update_step)
critic_lr = critic_learning_rate_decay_fn(update_step)
actor_lr = min(args.actor_lr, max(args.actor_lr_end, actor_lr))
critic_lr = min(args.critic_lr, max(args.critic_lr_end, critic_lr))
while True:
try:
replay = episodes_queue.get_nowait()
for (observation, action, reward, next_observation, done) in replay:
buffer.add(observation, action, reward, next_observation, done)
received_examples += len(replay)
except py_queue.Empty:
break
if len(buffer) >= args.train_steps:
if args.prioritized_replay:
beta = beta_deacy_fn(update_step)
beta = min(1.0, max(args.prioritized_replay_beta0, beta))
(tr_observations, tr_actions, tr_rewards, tr_next_observations, tr_dones,
weights, batch_idxes) = \
buffer.sample(
batch_size=args.batch_size,
beta=beta)
else:
(tr_observations, tr_actions, tr_rewards, tr_next_observations, tr_dones) = \
buffer.sample(batch_size=args.batch_size)
weights, batch_idxes = np.ones_like(tr_rewards), None
step_metrics, step_info = update_fn(
tr_observations, tr_actions, tr_rewards,
tr_next_observations, tr_dones,
weights, actor_lr, critic_lr)
update_step += 1
global_update_step.value += 1
if args.prioritized_replay:
new_priorities = np.abs(step_info["td_error"]) + 1e-6
buffer.update_priorities(batch_idxes, new_priorities)
for key, value in step_metrics.items():
value = to_numpy(value)[0]
logger.scalar_summary(key, value, update_step)
logger.scalar_summary("actor lr", actor_lr, update_step)
logger.scalar_summary("critic lr", critic_lr, update_step)
if update_step % args.save_step == 0:
save_fn(update_step)
else:
time.sleep(1)
logger.scalar_summary("buffer size", len(buffer), global_episode.value)
logger.scalar_summary(
"updates per example",
update_step * args.batch_size / received_examples,
global_episode.value)
save_fn(update_step)
raise KeyboardInterrupt
def play_single_thread(
actor, critic, target_actor, target_critic, args, prepare_fn,
global_episode, global_update_step, episodes_queue,
best_reward):
workerseed = args.seed + 241 * args.thread
set_global_seeds(workerseed)
args.logdir = "{}/thread_{}".format(args.logdir, args.thread)
create_if_need(args.logdir)
act_fn, _, save_fn = prepare_fn(actor, critic, target_actor, target_critic, args)
logger = Logger(args.logdir)
env = create_env(args)
random_process = create_random_process(args)
epsilon_cycle_len = random.randint(args.epsilon_cycle_len // 2, args.epsilon_cycle_len * 2)
epsilon_decay_fn = create_decay_fn(
"cycle",
initial_value=args.initial_epsilon,
final_value=args.final_epsilon,
cycle_len=epsilon_cycle_len,
num_cycles=args.max_episodes // epsilon_cycle_len)
episode = 1
step = 0
start_time = time.time()
while global_episode.value < args.max_episodes * (args.num_threads - args.num_train_threads) \
and global_update_step.value < args.max_update_steps * args.num_train_threads:
if episode % 100 == 0:
env = create_env(args)
seed = random.randrange(2 ** 32 - 2)
epsilon = min(args.initial_epsilon, max(args.final_epsilon, epsilon_decay_fn(episode)))
episode_metrics = {
"reward": 0.0,
"step": 0,
"epsilon": epsilon
}
observation = env.reset(seed=seed, difficulty=args.difficulty)
random_process.reset_states()
done = False
replay = []
while not done:
action = act_fn(observation, noise=epsilon * random_process.sample())
next_observation, reward, done, _ = env.step(action)
replay.append((observation, action, reward, next_observation, done))
episode_metrics["reward"] += reward
episode_metrics["step"] += 1
observation = next_observation
episodes_queue.put(replay)
episode += 1
global_episode.value += 1
if episode_metrics["reward"] > best_reward.value:
best_reward.value = episode_metrics["reward"]
logger.scalar_summary("best reward", best_reward.value, episode)
if episode_metrics["reward"] > 15.0 * args.reward_scale:
save_fn(episode)
step += episode_metrics["step"]
elapsed_time = time.time() - start_time
for key, value in episode_metrics.items():
logger.scalar_summary(key, value, episode)
logger.scalar_summary(
"episode per minute",
episode / elapsed_time * 60,
episode)
logger.scalar_summary(
"step per second",
step / elapsed_time,
episode)
if elapsed_time > 86400 * args.max_train_days:
global_episode.value = args.max_episodes * (args.num_threads - args.num_train_threads) + 1
raise KeyboardInterrupt
+90
View File
@@ -0,0 +1,90 @@
import numpy as np
import torch
import torch.nn as nn
from common.nets import LinearNet
from common.modules.NoisyLinear import NoisyLinear
def fanin_init(size, fanin=None):
fanin = fanin or size[0]
v = 1. / np.sqrt(fanin)
return torch.Tensor(size).uniform_(-v, v)
class Actor(nn.Module):
def __init__(self, n_observation, n_action,
layers, activation=torch.nn.ELU,
layer_norm=False,
parameters_noise=False, parameters_noise_factorised=False,
last_activation=torch.nn.Tanh, init_w=3e-3):
super(Actor, self).__init__()
if parameters_noise:
def linear_layer(x_in, x_out):
return NoisyLinear(x_in, x_out, factorised=parameters_noise_factorised)
else:
linear_layer = nn.Linear
self.feature_net = LinearNet(
layers=[n_observation] + layers,
activation=activation,
layer_norm=layer_norm,
linear_layer=linear_layer)
self.policy_net = LinearNet(
layers=[self.feature_net.output_shape, n_action],
activation=last_activation,
layer_norm=False
)
self.init_weights(init_w)
def init_weights(self, init_w):
for layer in self.feature_net.net:
if isinstance(layer, nn.Linear):
layer.weight.data = fanin_init(layer.weight.data.size())
for layer in self.feature_net.net:
if isinstance(layer, nn.Linear):
layer.weight.data.uniform_(-init_w, init_w)
def forward(self, observation):
x = observation
x = self.feature_net.forward(x)
x = self.policy_net.forward(x)
return x
class Critic(nn.Module):
def __init__(self, n_observation, n_action,
layers, activation=torch.nn.ELU,
layer_norm=False,
parameters_noise=False, parameters_noise_factorised=False,
init_w=3e-3):
super(Critic, self).__init__()
if parameters_noise:
def linear_layer(x_in, x_out):
return NoisyLinear(x_in, x_out, factorised=parameters_noise_factorised)
else:
linear_layer = nn.Linear
self.feature_net = LinearNet(
layers=[n_observation + n_action] + layers,
activation=activation,
layer_norm=layer_norm,
linear_layer=linear_layer)
self.value_net = nn.Linear(self.feature_net.output_shape, 1)
self.init_weights(init_w)
def init_weights(self, init_w):
for layer in self.feature_net.net:
if isinstance(layer, nn.Linear):
layer.weight.data = fanin_init(layer.weight.data.size())
self.value_net.weight.data.uniform_(-init_w, init_w)
def forward(self, observation, action):
x = torch.cat((observation, action), dim=1)
x = self.feature_net.forward(x)
x = self.value_net.forward(x)
return x
+186
View File
@@ -0,0 +1,186 @@
import os
import json
import argparse
import numpy as np
import pandas as pd
import torch
from pprint import pprint
from osim.env import RunEnv
from osim.http.client import Client
from common.misc_util import boolean_flag, query_yes_no
from common.env_wrappers import create_observation_handler, create_action_handler, create_env
from ddpg.train import str2params, activations
from ddpg.model import create_model, create_act_update_fns
REMOTE_BASE = 'http://grader.crowdai.org:1729'
ACTION_SHAPE = 18
SEEDS = [
3834825972, 3049289152, 3538742899, 2904257823, 4011088434,
2684066875, 781202090, 1691535473, 898088606, 1301477286
]
def parse_args():
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
parser.add_argument('--restore-args-from', type=str, default=None)
parser.add_argument('--restore-actor-from', type=str, default=None)
parser.add_argument('--restore-critic-from', type=str, default=None)
parser.add_argument('--max-obstacles', type=int, default=3)
parser.add_argument('--num-episodes', type=int, default=1)
parser.add_argument('--token', type=str, default=None)
boolean_flag(parser, "visualize", default=False)
boolean_flag(parser, "submit", default=False)
return parser.parse_args()
def restore_args(args):
with open(args.restore_args_from, "r") as fin:
params = json.load(fin)
unwanted = [
"max_obstacles",
"restore_args_from",
"restore_actor_from",
"restore_critic_from"
]
for unwanted_key in unwanted:
value = params.pop(unwanted_key, None)
if value is not None:
del value
for key, value in params.items():
setattr(args, key, value)
return args
def submit(actor, critic, args, act_update_fn):
act_fn, _, _ = act_update_fn(actor, critic, None, None, args)
client = Client(REMOTE_BASE)
all_episode_metrics = []
episode_metrics = {
"reward": 0.0,
"step": 0,
}
observation_handler = create_observation_handler(args)
action_handler = create_action_handler(args)
observation = client.env_create(args.token)
action = np.zeros(ACTION_SHAPE, dtype=np.float32)
observation = observation_handler(observation, action)
submitted = False
while not submitted:
print(episode_metrics["reward"])
action = act_fn(observation)
observation, reward, done, _ = client.env_step(action_handler(action).tolist())
episode_metrics["reward"] += reward
episode_metrics["step"] += 1
if done:
all_episode_metrics.append(episode_metrics)
episode_metrics = {
"reward": 0.0,
"step": 0,
}
observation_handler = create_observation_handler(args)
action_handler = create_action_handler(args)
observation = client.env_create(args.token)
if not observation:
submitted = True
break
action = np.zeros(ACTION_SHAPE, dtype=np.float32)
observation = observation_handler(observation, action)
else:
observation = observation_handler(observation, action)
df = pd.DataFrame(all_episode_metrics)
pprint(df.describe())
if query_yes_no("Submit?"):
client.submit()
def test(actor, critic, args, act_update_fn):
act_fn, _, _ = act_update_fn(actor, critic, None, None, args)
env = RunEnv(visualize=args.visualize, max_obstacles=args.max_obstacles)
all_episode_metrics = []
for episode in range(args.num_episodes):
episode_metrics = {
"reward": 0.0,
"step": 0,
}
observation_handler = create_observation_handler(args)
action_handler = create_action_handler(args)
observation = env.reset(difficulty=2, seed=SEEDS[episode % len(SEEDS)])
action = np.zeros(ACTION_SHAPE, dtype=np.float32)
observation = observation_handler(observation, action)
done = False
while not done:
print(episode_metrics["reward"])
action = act_fn(observation)
observation, reward, done, _ = env.step(action_handler(action))
episode_metrics["reward"] += reward
episode_metrics["step"] += 1
if done:
break
observation = observation_handler(observation, action)
all_episode_metrics.append(episode_metrics)
df = pd.DataFrame(all_episode_metrics)
pprint(df.describe())
def submit_or_test(args, model_fn, act_update_fn, submit_fn, test_fn):
args = restore_args(args)
env = create_env(args)
args.n_action = env.action_space.shape[0]
args.n_observation = env.observation_space.shape[0]
args.actor_layers = str2params(args.actor_layers)
args.critic_layers = str2params(args.critic_layers)
args.actor_activation = activations[args.actor_activation]
args.critic_activation = activations[args.critic_activation]
actor, critic = model_fn(args)
actor.load_state_dict(torch.load(args.restore_actor_from))
critic.load_state_dict(torch.load(args.restore_critic_from))
if args.submit:
submit_fn(actor, critic, args, act_update_fn)
else:
test_fn(actor, critic, args, act_update_fn)
if __name__ == '__main__':
os.environ['OMP_NUM_THREADS'] = '1'
torch.set_num_threads(1)
args = parse_args()
submit_or_test(args, create_model, create_act_update_fns, submit, test)
+237
View File
@@ -0,0 +1,237 @@
import argparse
import os
import json
import copy
import torch
import torch.multiprocessing as mp
from multiprocessing import Value
from common.misc_util import boolean_flag, str2params, create_if_need
from common.env_wrappers import create_env
from common.torch_util import activations, hard_update
from ddpg.model import create_model, create_act_update_fns, train_multi_thread, \
train_single_thread, play_single_thread
def parse_args():
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
parser.add_argument('--seed', type=int, default=42)
parser.add_argument('--difficulty', type=int, default=2)
parser.add_argument('--max-obstacles', type=int, default=3)
parser.add_argument('--logdir', type=str, default="./logs")
parser.add_argument('--num-threads', type=int, default=1)
parser.add_argument('--num-train-threads', type=int, default=1)
boolean_flag(parser, "ddpg-wrapper", default=False)
parser.add_argument('--skip-frames', type=int, default=1)
parser.add_argument('--fail-reward', type=float, default=0.0)
parser.add_argument('--reward-scale', type=float, default=1.)
boolean_flag(parser, "flip-state-action", default=False)
for agent in ["actor", "critic"]:
parser.add_argument('--{}-layers'.format(agent), type=str, default="64-64")
parser.add_argument('--{}-activation'.format(agent), type=str, default="relu")
boolean_flag(parser, "{}-layer-norm".format(agent), default=False)
boolean_flag(parser, "{}-parameters-noise".format(agent), default=False)
boolean_flag(parser, "{}-parameters-noise-factorised".format(agent), default=False)
parser.add_argument('--{}-lr'.format(agent), type=float, default=1e-3)
parser.add_argument('--{}-lr-end'.format(agent), type=float, default=5e-5)
parser.add_argument('--restore-{}-from'.format(agent), type=str, default=None)
parser.add_argument('--gamma', type=float, default=0.96)
parser.add_argument('--loss-type', type=str, default="quadric-linear")
parser.add_argument('--grad-clip', type=float, default=10.)
parser.add_argument('--tau', default=0.01, type=float)
parser.add_argument('--train-steps', type=int, default=int(1e4))
parser.add_argument('--batch-size', type=int, default=256) # per worker
parser.add_argument('--buffer-size', type=int, default=int(1e6))
boolean_flag(parser, "prioritized-replay", default=False)
parser.add_argument('--prioritized-replay-alpha', default=0.6, type=float)
parser.add_argument('--prioritized-replay-beta0', default=0.4, type=float)
parser.add_argument('--initial-epsilon', default=1., type=float)
parser.add_argument('--final-epsilon', default=0.01, type=float)
parser.add_argument('--max-episodes', default=int(1e4), type=int)
parser.add_argument('--max-update-steps', default=int(5e6), type=int)
parser.add_argument('--epsilon-cycle-len', default=int(2e2), type=int)
parser.add_argument('--max-train-days', default=int(1e1), type=int)
parser.add_argument('--rp-type', default="ornstein-uhlenbeck", type=str)
parser.add_argument('--rp-theta', default=0.15, type=float)
parser.add_argument('--rp-sigma', default=0.2, type=float)
parser.add_argument('--rp-sigma-min', default=0.15, type=float)
parser.add_argument('--rp-mu', default=0.0, type=float)
parser.add_argument('--clip-delta', type=int, default=10)
parser.add_argument('--save-step', type=int, default=int(1e4))
parser.add_argument('--restore-args-from', type=str, default=None)
return parser.parse_args()
def restore_args(args):
with open(args.restore_args_from, "r") as fin:
params = json.load(fin)
del params["seed"]
del params["difficulty"]
del params["max_obstacles"]
del params["logdir"]
del params["num_threads"]
del params["num_train_threads"]
del params["skip_frames"]
for agent in ["actor", "critic"]:
del params["{}_lr".format(agent)]
del params["{}_lr_end".format(agent)]
del params["restore_{}_from".format(agent)]
del params["grad_clip"]
del params["tau"]
del params["train_steps"]
del params["batch_size"]
del params["buffer_size"]
del params["prioritized_replay"]
del params["prioritized_replay_alpha"]
del params["prioritized_replay_beta0"]
del params["initial_epsilon"]
del params["final_epsilon"]
del params["max_episodes"]
del params["max_update_steps"]
del params["epsilon_cycle_len"]
del params["max_train_days"]
del params["rp_type"]
del params["rp_theta"]
del params["rp_sigma"]
del params["rp_sigma_min"]
del params["rp_mu"]
del params["clip_delta"]
del params["save_step"]
del params["restore_args_from"]
for key, value in params.items():
setattr(args, key, value)
return args
def train(args, model_fn, act_update_fns, multi_thread, train_single, play_single):
create_if_need(args.logdir)
if args.restore_args_from is not None:
args = restore_args(args)
with open("{}/args.json".format(args.logdir), "w") as fout:
json.dump(vars(args), fout, indent=4, ensure_ascii=False, sort_keys=True)
env = create_env(args)
if args.flip_state_action and hasattr(env, "state_transform"):
args.flip_states = env.state_transform.flip_states
args.batch_size = args.batch_size // 2
args.n_action = env.action_space.shape[0]
args.n_observation = env.observation_space.shape[0]
args.actor_layers = str2params(args.actor_layers)
args.critic_layers = str2params(args.critic_layers)
args.actor_activation = activations[args.actor_activation]
args.critic_activation = activations[args.critic_activation]
actor, critic = model_fn(args)
if args.restore_actor_from is not None:
actor.load_state_dict(torch.load(args.restore_actor_from))
if args.restore_critic_from is not None:
critic.load_state_dict(torch.load(args.restore_critic_from))
actor.train()
critic.train()
actor.share_memory()
critic.share_memory()
target_actor = copy.deepcopy(actor)
target_critic = copy.deepcopy(critic)
hard_update(target_actor, actor)
hard_update(target_critic, critic)
target_actor.train()
target_critic.train()
target_actor.share_memory()
target_critic.share_memory()
_, _, save_fn = act_update_fns(actor, critic, target_actor, target_critic, args)
processes = []
best_reward = Value("f", 0.0)
try:
if args.num_threads == args.num_train_threads:
for rank in range(args.num_threads):
args.thread = rank
p = mp.Process(
target=multi_thread,
args=(actor, critic, target_actor, target_critic, args, act_update_fns,
best_reward))
p.start()
processes.append(p)
else:
global_episode = Value("i", 0)
global_update_step = Value("i", 0)
episodes_queue = mp.Queue()
for rank in range(args.num_threads):
args.thread = rank
if rank < args.num_train_threads:
p = mp.Process(
target=train_single,
args=(actor, critic, target_actor, target_critic, args, act_update_fns,
global_episode, global_update_step, episodes_queue))
else:
p = mp.Process(
target=play_single,
args=(actor, critic, target_actor, target_critic, args, act_update_fns,
global_episode, global_update_step, episodes_queue,
best_reward))
p.start()
processes.append(p)
for p in processes:
p.join()
except KeyboardInterrupt:
pass
save_fn()
if __name__ == '__main__':
os.environ['OMP_NUM_THREADS'] = '1'
torch.set_num_threads(1)
args = parse_args()
train(args,
create_model,
create_act_update_fns,
train_multi_thread,
train_single_thread,
play_single_thread)