mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Use openai atari wrapper
This commit is contained in:
+2
-1
@@ -52,6 +52,7 @@ class A2CAgent:
|
||||
rollout = []
|
||||
states = self.states
|
||||
for i in range(config.rollout_length):
|
||||
states = self.task.normalize_state(states)
|
||||
prob, log_prob, value = self.network.predict(states)
|
||||
actions = [self.policy.sample(p) for p in prob.data.cpu().numpy()]
|
||||
actions = config.action_shift_fn(actions)
|
||||
@@ -68,7 +69,7 @@ class A2CAgent:
|
||||
states = next_states
|
||||
|
||||
self.states = states
|
||||
_, _, pending_value = self.network.predict(states)
|
||||
_, _, pending_value = self.network.predict(self.task.normalize_state(states))
|
||||
rollout.append([None, None, pending_value, None, None, None])
|
||||
|
||||
processed_rollout = [None] * (len(rollout) - 1)
|
||||
|
||||
@@ -40,7 +40,7 @@ class NStepDQNAgent:
|
||||
rollout = []
|
||||
states = self.states
|
||||
for i in range(config.rollout_length):
|
||||
q = self.learning_network.predict(states)
|
||||
q = self.learning_network.predict(self.task.normalize_state(states))
|
||||
actions = [self.policy.sample(v) for v in q.data.cpu().numpy()]
|
||||
actions = config.action_shift_fn(actions)
|
||||
next_states, rewards, terminals, _ = self.task.step(actions)
|
||||
@@ -63,7 +63,7 @@ class NStepDQNAgent:
|
||||
self.states = states
|
||||
|
||||
processed_rollout = [None] * (len(rollout))
|
||||
returns = self.target_network.predict(states).data
|
||||
returns = self.target_network.predict(self.task.normalize_state(states)).data
|
||||
returns, _ = torch.max(returns, dim=1, keepdim=True)
|
||||
for i in reversed(range(len(rollout))):
|
||||
q, actions, rewards, terminals = rollout[i]
|
||||
|
||||
+1
-1
@@ -4,4 +4,4 @@ from .DDPG_agent import *
|
||||
from .A2C_agent import *
|
||||
from .CategoricalDQN_agent import *
|
||||
from .NStepDQN_agent import *
|
||||
from .QuantileRegressionDQN_agent import *
|
||||
from .QuantileRegressionDQN_agent import *
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from .atari_wrapper import *
|
||||
from .atari_wrapper import *
|
||||
from .policy import *
|
||||
from .replay import *
|
||||
from .task import *
|
||||
|
||||
+157
-86
@@ -1,54 +1,70 @@
|
||||
# based on https://github.com/openai/baselines/blob/master/baselines/common/atari_wrappers.py
|
||||
|
||||
import numpy as np
|
||||
from collections import deque
|
||||
import gym
|
||||
from gym import spaces
|
||||
from skimage import color, transform
|
||||
from gym.spaces import Box
|
||||
import cv2
|
||||
cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
class NoopResetEnv(gym.Wrapper):
|
||||
def __init__(self, env=None, noop_max=30):
|
||||
def __init__(self, env, noop_max=30):
|
||||
"""Sample initial states by taking random number of no-ops on reset.
|
||||
No-op is assumed to be action 0.
|
||||
"""
|
||||
super(NoopResetEnv, self).__init__(env)
|
||||
gym.Wrapper.__init__(self, env)
|
||||
self.noop_max = noop_max
|
||||
self.override_num_noops = None
|
||||
self.noop_action = 0
|
||||
assert env.unwrapped.get_action_meanings()[0] == 'NOOP'
|
||||
|
||||
def reset(self):
|
||||
def reset(self, **kwargs):
|
||||
""" Do no-op action for a number of steps in [1, noop_max]."""
|
||||
self.env.reset()
|
||||
noops = np.random.randint(1, self.noop_max + 1)
|
||||
self.env.reset(**kwargs)
|
||||
if self.override_num_noops is not None:
|
||||
noops = self.override_num_noops
|
||||
else:
|
||||
noops = self.unwrapped.np_random.randint(1, self.noop_max + 1) #pylint: disable=E1101
|
||||
assert noops > 0
|
||||
obs = None
|
||||
for _ in range(noops):
|
||||
obs, _, _, _ = self.env.step(0)
|
||||
obs, _, done, _ = self.env.step(self.noop_action)
|
||||
if done:
|
||||
obs = self.env.reset(**kwargs)
|
||||
return obs
|
||||
|
||||
def step(self, action):
|
||||
return self.env.step(action)
|
||||
def step(self, ac):
|
||||
return self.env.step(ac)
|
||||
|
||||
class FireResetEnv(gym.Wrapper):
|
||||
def __init__(self, env=None):
|
||||
def __init__(self, env):
|
||||
"""Take action on reset for environments that are fixed until firing."""
|
||||
super(FireResetEnv, self).__init__(env)
|
||||
gym.Wrapper.__init__(self, env)
|
||||
assert env.unwrapped.get_action_meanings()[1] == 'FIRE'
|
||||
assert len(env.unwrapped.get_action_meanings()) >= 3
|
||||
|
||||
def reset(self):
|
||||
self.env.reset()
|
||||
obs, _, _, _ = self.env.step(1)
|
||||
obs, _, _, _ = self.env.step(2)
|
||||
def reset(self, **kwargs):
|
||||
self.env.reset(**kwargs)
|
||||
obs, _, done, _ = self.env.step(1)
|
||||
if done:
|
||||
self.env.reset(**kwargs)
|
||||
obs, _, done, _ = self.env.step(2)
|
||||
if done:
|
||||
self.env.reset(**kwargs)
|
||||
return obs
|
||||
|
||||
def step(self, action):
|
||||
return self.env.step(action)
|
||||
def step(self, ac):
|
||||
return self.env.step(ac)
|
||||
|
||||
class EpisodicLifeEnv(gym.Wrapper):
|
||||
def __init__(self, env=None):
|
||||
def __init__(self, env):
|
||||
"""Make end-of-life == end-of-episode, but only reset on true game over.
|
||||
Done by DeepMind for the DQN and co. since it helps value estimation.
|
||||
"""
|
||||
super(EpisodicLifeEnv, self).__init__(env)
|
||||
gym.Wrapper.__init__(self, env)
|
||||
self.lives = 0
|
||||
self.was_real_done = True
|
||||
self.was_realreset = False
|
||||
|
||||
def step(self, action):
|
||||
obs, reward, done, info = self.env.step(action)
|
||||
@@ -57,56 +73,53 @@ class EpisodicLifeEnv(gym.Wrapper):
|
||||
# then update lives to handle bonus lives
|
||||
lives = self.env.unwrapped.ale.lives()
|
||||
if lives < self.lives and lives > 0:
|
||||
# for Qbert somtimes we stay in lives == 0 condtion for a few frames
|
||||
# for Qbert sometimes we stay in lives == 0 condtion for a few frames
|
||||
# so its important to keep lives > 0, so that we only reset once
|
||||
# the environment advertises done.
|
||||
done = True
|
||||
self.lives = lives
|
||||
return obs, reward, done, info
|
||||
|
||||
def reset(self):
|
||||
def reset(self, **kwargs):
|
||||
"""Reset only when lives are exhausted.
|
||||
This way all states are still reachable even though lives are episodic,
|
||||
and the learner need not know about any of this behind-the-scenes.
|
||||
"""
|
||||
if self.was_real_done:
|
||||
obs = self.env.reset()
|
||||
self.was_realreset = True
|
||||
obs = self.env.reset(**kwargs)
|
||||
else:
|
||||
# no-op step to advance from terminal/lost life state
|
||||
obs, _, _, _ = self.env.step(0)
|
||||
self.was_realreset = False
|
||||
self.lives = self.env.unwrapped.ale.lives()
|
||||
return obs
|
||||
|
||||
class MaxAndSkipEnv(gym.Wrapper):
|
||||
def __init__(self, env=None, skip=4):
|
||||
def __init__(self, env, skip=4):
|
||||
"""Return only every `skip`-th frame"""
|
||||
super(MaxAndSkipEnv, self).__init__(env)
|
||||
gym.Wrapper.__init__(self, env)
|
||||
# most recent raw observations (for max pooling across time steps)
|
||||
self._obs_buffer = deque(maxlen=2)
|
||||
self._obs_buffer = np.zeros((2,)+env.observation_space.shape, dtype=np.uint8)
|
||||
self._skip = skip
|
||||
|
||||
def step(self, action):
|
||||
"""Repeat action, sum reward, and max over last observations."""
|
||||
total_reward = 0.0
|
||||
done = None
|
||||
for _ in range(self._skip):
|
||||
for i in range(self._skip):
|
||||
obs, reward, done, info = self.env.step(action)
|
||||
self._obs_buffer.append(obs)
|
||||
if i == self._skip - 2: self._obs_buffer[0] = obs
|
||||
if i == self._skip - 1: self._obs_buffer[1] = obs
|
||||
total_reward += reward
|
||||
if done:
|
||||
break
|
||||
|
||||
max_frame = np.max(np.stack(self._obs_buffer), axis=0)
|
||||
# Note that the observation on the done=True frame
|
||||
# doesn't matter
|
||||
max_frame = self._obs_buffer.max(axis=0)
|
||||
|
||||
return max_frame, total_reward, done, info
|
||||
|
||||
def reset(self):
|
||||
"""Clear past frame buffer and init. to first obs. from inner env."""
|
||||
self._obs_buffer.clear()
|
||||
obs = self.env.reset()
|
||||
self._obs_buffer.append(obs)
|
||||
return obs
|
||||
def reset(self, **kwargs):
|
||||
return self.env.reset(**kwargs)
|
||||
|
||||
class SkipEnv(gym.Wrapper):
|
||||
def __init__(self, env=None, skip=4):
|
||||
@@ -129,6 +142,92 @@ class SkipEnv(gym.Wrapper):
|
||||
obs = self.env.reset()
|
||||
return obs
|
||||
|
||||
class ClipRewardEnv(gym.RewardWrapper):
|
||||
def __init__(self, env):
|
||||
gym.RewardWrapper.__init__(self, env)
|
||||
|
||||
def reward(self, reward):
|
||||
"""Bin reward to {+1, 0, -1} by its sign."""
|
||||
return np.sign(reward)
|
||||
|
||||
class WarpFrame(gym.ObservationWrapper):
|
||||
def __init__(self, env):
|
||||
"""Warp frames to 84x84 as done in the Nature paper and later work."""
|
||||
gym.ObservationWrapper.__init__(self, env)
|
||||
self.width = 84
|
||||
self.height = 84
|
||||
self.observation_space = spaces.Box(low=0, high=255,
|
||||
shape=(self.height, self.width, 1), dtype=np.uint8)
|
||||
|
||||
def observation(self, frame):
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY)
|
||||
frame = cv2.resize(frame, (self.width, self.height), interpolation=cv2.INTER_AREA)
|
||||
return frame[:, :, None]
|
||||
|
||||
class LazyFrames(object):
|
||||
def __init__(self, frames):
|
||||
"""This object ensures that common frames between the observations are only stored once.
|
||||
It exists purely to optimize memory usage which can be huge for DQN's 1M frames replay
|
||||
buffers.
|
||||
|
||||
This object should only be converted to numpy array before being passed to the model.
|
||||
|
||||
You'd not believe how complex the previous solution was."""
|
||||
self._frames = frames
|
||||
self._out = None
|
||||
|
||||
def _force(self):
|
||||
if self._out is None:
|
||||
self._out = np.concatenate(self._frames, axis=2)
|
||||
self._frames = None
|
||||
return self._out
|
||||
|
||||
def __array__(self, dtype=None):
|
||||
out = self._force()
|
||||
if dtype is not None:
|
||||
out = out.astype(dtype)
|
||||
return out
|
||||
|
||||
def __len__(self):
|
||||
return len(self._force())
|
||||
|
||||
def __getitem__(self, i):
|
||||
return self._force()[i]
|
||||
|
||||
class StackFrame(gym.Wrapper):
|
||||
def __init__(self, env=None, history_length=1):
|
||||
super(StackFrame, self).__init__(env)
|
||||
self.history_length = history_length
|
||||
self.buffer = None
|
||||
|
||||
def reset(self):
|
||||
state = self.env.reset()
|
||||
self.buffer = [state] * self.history_length
|
||||
# return LazyFrames(self.buffer)
|
||||
return np.asarray(np.vstack(self.buffer))
|
||||
|
||||
def step(self, action):
|
||||
state, reward, done, info = self.env.step(action)
|
||||
self.buffer.pop(0)
|
||||
self.buffer.append(state)
|
||||
# return LazyFrames(self.buffer), reward, done, info
|
||||
return np.asarray(np.vstack(self.buffer)), reward, done, info
|
||||
|
||||
class WrapPyTorch(gym.ObservationWrapper):
|
||||
# from https://github.com/ikostrikov/pytorch-a2c-ppo-acktr/blob/master/envs.py
|
||||
def __init__(self, env=None):
|
||||
super(WrapPyTorch, self).__init__(env)
|
||||
obs_shape = self.observation_space.shape
|
||||
self.observation_space = Box(
|
||||
self.observation_space.low[0,0,0],
|
||||
self.observation_space.high[0,0,0],
|
||||
[obs_shape[2], obs_shape[1], obs_shape[0]],
|
||||
dtype=np.uint8
|
||||
)
|
||||
|
||||
def observation(self, observation):
|
||||
return observation.transpose(2, 0, 1)
|
||||
|
||||
class DatasetEnv(gym.Wrapper):
|
||||
def __init__(self, env=None):
|
||||
super(DatasetEnv, self).__init__(env)
|
||||
@@ -153,52 +252,24 @@ class DatasetEnv(gym.Wrapper):
|
||||
self.saved_obs.append(obs)
|
||||
return obs
|
||||
|
||||
class ProcessFrame(gym.Wrapper):
|
||||
def __init__(self, env=None, frame_size=84):
|
||||
super(ProcessFrame, self).__init__(env)
|
||||
self.frame_size = frame_size
|
||||
self.observation_space = spaces.Box(low=0, high=255, shape=(1, frame_size, frame_size), dtype=np.uint8)
|
||||
def make_atari(env_id, frame_skip=4):
|
||||
env = gym.make(env_id)
|
||||
assert 'NoFrameskip' in env.spec.id
|
||||
env = NoopResetEnv(env, noop_max=30)
|
||||
env = MaxAndSkipEnv(env, skip=4)
|
||||
return env
|
||||
|
||||
def process(self, obs):
|
||||
obs = color.rgb2gray(obs)
|
||||
obs = transform.resize(obs, (self.frame_size, self.frame_size), mode='constant')
|
||||
obs = (255 * obs).astype(np.uint8).reshape((1, ) + obs.shape)
|
||||
return obs
|
||||
|
||||
def step(self, action):
|
||||
obs, reward, done, info = self.env.step(action)
|
||||
return self.process(obs), reward, done, info
|
||||
|
||||
def reset(self):
|
||||
return self.process(self.env.reset())
|
||||
|
||||
class NormalizeFrame(gym.Wrapper):
|
||||
def __init__(self, env=None):
|
||||
super(NormalizeFrame, self).__init__(env)
|
||||
|
||||
def _normalize(self, obs):
|
||||
return np.asarray(obs, dtype=np.float32) / 255.0
|
||||
|
||||
def step(self, action):
|
||||
obs, reward, done, info = self.env.step(action)
|
||||
return self._normalize(obs), reward, done, info
|
||||
|
||||
def reset(self):
|
||||
return self._normalize(self.env.reset())
|
||||
|
||||
class StackFrame(gym.Wrapper):
|
||||
def __init__(self, env=None, history_length=1):
|
||||
super(StackFrame, self).__init__(env)
|
||||
self.history_length = history_length
|
||||
self.buffer = None
|
||||
|
||||
def reset(self):
|
||||
state = self.env.reset()
|
||||
self.buffer = [state] * self.history_length
|
||||
return np.asarray(np.vstack(self.buffer))
|
||||
|
||||
def step(self, action):
|
||||
state, reward, done, info = self.env.step(action)
|
||||
self.buffer.pop(0)
|
||||
self.buffer.append(state)
|
||||
return np.asarray(np.vstack(self.buffer)), reward, done, info
|
||||
def wrap_deepmind(env, episode_life=True, clip_rewards=True, history_length=0):
|
||||
"""Configure environment for DeepMind-style Atari.
|
||||
"""
|
||||
if episode_life:
|
||||
env = EpisodicLifeEnv(env)
|
||||
if 'FIRE' in env.unwrapped.get_action_meanings():
|
||||
env = FireResetEnv(env)
|
||||
env = WarpFrame(env)
|
||||
if clip_rewards:
|
||||
env = ClipRewardEnv(env)
|
||||
env = WrapPyTorch(env)
|
||||
if history_length:
|
||||
env = StackFrame(env, history_length)
|
||||
return env
|
||||
|
||||
+49
-5
@@ -225,12 +225,17 @@ class GeneralReplay:
|
||||
|
||||
def feed(self, experiences):
|
||||
for experience in zip(*experiences):
|
||||
self.buffer.append(experience)
|
||||
if len(self.buffer) > self.memory_size:
|
||||
del self.buffer[0]
|
||||
self.feed_single(experience)
|
||||
|
||||
def sample(self):
|
||||
sampled = zip(*random.sample(self.buffer, self.batch_size))
|
||||
def feed_single(self, experience):
|
||||
self.buffer.append(experience)
|
||||
if len(self.buffer) > self.memory_size:
|
||||
del self.buffer[0]
|
||||
|
||||
def sample(self, batch_size=None):
|
||||
if batch_size is None:
|
||||
batch_size = self.batch_size
|
||||
sampled = zip(*random.sample(self.buffer, batch_size))
|
||||
return sampled
|
||||
|
||||
def clear(self):
|
||||
@@ -238,3 +243,42 @@ class GeneralReplay:
|
||||
|
||||
def full(self):
|
||||
return len(self.buffer) == self.memory_size
|
||||
|
||||
def size(self):
|
||||
return len(self.buffer)
|
||||
|
||||
def empty(self):
|
||||
return not len(self.buffer)
|
||||
|
||||
class SkewedReplay:
|
||||
def __init__(self, memory_size, batch_size):
|
||||
memory_size = memory_size / 2
|
||||
self.non_zero_reward = GeneralReplay(memory_size, batch_size / 2)
|
||||
self.zero_reward = GeneralReplay(memory_size, batch_size / 2)
|
||||
self.batch_size = batch_size
|
||||
|
||||
def feed(self, experiences):
|
||||
experiences = zip(*experiences)
|
||||
for exp in experiences:
|
||||
if np.abs(exp[2]) < 1e-5:
|
||||
self.zero_reward.feed_single(exp)
|
||||
else:
|
||||
self.non_zero_reward.feed_single(exp)
|
||||
|
||||
def sample(self):
|
||||
if self.zero_reward.empty():
|
||||
batch = self.non_zero_reward.sample(self.batch_size)
|
||||
elif self.non_zero_reward.empty():
|
||||
batch = self.zero_reward.sample(self.batch_size)
|
||||
else:
|
||||
non_zero_batch_size = min(self.non_zero_reward.size(), self.batch_size / 2)
|
||||
zero_batch_size = min(self.zero_reward.size(), self.batch_size / 2)
|
||||
batch1 = self.zero_reward.sample(zero_batch_size)
|
||||
batch2 = self.non_zero_reward.sample(non_zero_batch_size)
|
||||
batch = list(map(lambda seq: np.concatenate([np.asarray(x) for x in seq], axis=0), zip(batch1, batch2)))
|
||||
batch = list(map(lambda x: np.asarray(x), batch))
|
||||
return batch
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
+28
-29
@@ -11,6 +11,8 @@ import multiprocessing as mp
|
||||
import sys
|
||||
from .bench import Monitor
|
||||
from utils import *
|
||||
import datetime
|
||||
import uuid
|
||||
|
||||
class BasicTask:
|
||||
def __init__(self, max_steps=sys.maxsize):
|
||||
@@ -34,9 +36,6 @@ class BasicTask:
|
||||
def random_action(self):
|
||||
return self.env.action_space.sample()
|
||||
|
||||
def set_monitor(self, filename):
|
||||
self.env = Monitor(self.env, filename)
|
||||
|
||||
class ClassicalControl(BasicTask):
|
||||
def __init__(self, name='CartPole-v0', max_steps=200):
|
||||
BasicTask.__init__(self, max_steps)
|
||||
@@ -57,23 +56,22 @@ class LunarLander(BasicTask):
|
||||
self.state_dim = self.env.observation_space.shape[0]
|
||||
|
||||
class PixelAtari(BasicTask):
|
||||
def __init__(self, name, no_op, frame_skip, normalized_state=True,
|
||||
frame_size=84, max_steps=10000, history_length=1):
|
||||
def __init__(self, name, seed=0, log_file=None, max_steps=sys.maxsize,
|
||||
frame_skip=4, history_length=4):
|
||||
BasicTask.__init__(self, max_steps)
|
||||
self.normalized_state = normalized_state
|
||||
self.name = name
|
||||
env = gym.make(name)
|
||||
assert 'NoFrameskip' in env.spec.id
|
||||
env = EpisodicLifeEnv(env)
|
||||
env = NoopResetEnv(env, noop_max=no_op)
|
||||
env = MaxAndSkipEnv(env, skip=frame_skip)
|
||||
if 'FIRE' in env.unwrapped.get_action_meanings():
|
||||
env = FireResetEnv(env)
|
||||
env = ProcessFrame(env, frame_size)
|
||||
if normalized_state:
|
||||
env = NormalizeFrame(env)
|
||||
self.env = StackFrame(env, history_length)
|
||||
env = make_atari(name, frame_skip)
|
||||
env.seed(seed)
|
||||
if log_file is None:
|
||||
log_dir = '%s-%s' % (
|
||||
name,
|
||||
datetime.datetime.now().strftime("%y%m%d-%-H%M%S"))
|
||||
mkdir('./log/%s' % log_dir)
|
||||
log_file = './log/%s/%s' % (log_dir, uuid.uuid1())
|
||||
env = Monitor(env, log_file)
|
||||
env = wrap_deepmind(env, history_length=history_length)
|
||||
self.env = env
|
||||
self.action_dim = self.env.action_space.n
|
||||
self.name = name
|
||||
|
||||
def normalize_state(self, state):
|
||||
return np.asarray(state) / 255.0
|
||||
@@ -143,12 +141,12 @@ class Roboschool(BasicTask):
|
||||
def step(self, action):
|
||||
return BasicTask.step(self, np.clip(action, -1, 1))
|
||||
|
||||
def sub_task(parent_pipe, pipe, task_fn, filename=None):
|
||||
def sub_task(parent_pipe, pipe, task_fn, rank, log_dir):
|
||||
np.random.seed()
|
||||
seed = np.random.randint(0, sys.maxsize)
|
||||
parent_pipe.close()
|
||||
task = task_fn()
|
||||
if filename is not None:
|
||||
task.set_monitor(filename)
|
||||
task.env.seed(np.random.randint(0, sys.maxsize))
|
||||
task = task_fn(log_file=os.path.join(log_dir, str(rank)))
|
||||
task.env.seed(seed)
|
||||
while True:
|
||||
op, data = pipe.recv()
|
||||
if op == 'step':
|
||||
@@ -166,13 +164,11 @@ class ParallelizedTask:
|
||||
self.task_fn = task_fn
|
||||
self.task = task_fn()
|
||||
self.name = self.task.name
|
||||
# date = datetime.datetime.now().strftime("%I:%M%p-on-%B-%d-%Y")
|
||||
mkdir('./log/%s-%s' % (self.name, tag))
|
||||
filenames = ['./log/%s-%s/worker-%d' % (self.name, tag, i)
|
||||
for i in range(num_workers)]
|
||||
log_dir = './log/%s-%s' % (self.name, tag)
|
||||
mkdir(log_dir)
|
||||
self.pipes, worker_pipes = zip(*[mp.Pipe() for _ in range(num_workers)])
|
||||
args = [(p, wp, task_fn, filename)
|
||||
for p, wp, filename in zip(self.pipes, worker_pipes, filenames)]
|
||||
args = [(p, wp, task_fn, rank, log_dir)
|
||||
for rank, (p, wp) in enumerate(zip(self.pipes, worker_pipes))]
|
||||
self.workers = [mp.Process(target=sub_task, args=arg) for arg in args]
|
||||
for p in self.workers: p.start()
|
||||
for p in worker_pipes: p.close()
|
||||
@@ -200,3 +196,6 @@ class ParallelizedTask:
|
||||
for pipe in self.pipes:
|
||||
pipe.send(('exit', None))
|
||||
for p in self.workers: p.join()
|
||||
|
||||
def normalize_state(self, state):
|
||||
return self.task.normalize_state(state)
|
||||
|
||||
@@ -29,54 +29,11 @@ def dqn_cart_pole():
|
||||
# config.double_q = False
|
||||
run_episodes(DQNAgent(config))
|
||||
|
||||
def async_cart_pole():
|
||||
config = Config()
|
||||
config.task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200)
|
||||
config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.network_fn = lambda: FCNet([4, 50, 200, 2])
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=0.5, final_step=5000, min_epsilon=0.1)
|
||||
# config.worker = OneStepQLearning
|
||||
config.worker = NStepQLearning
|
||||
# config.worker = OneStepSarsa
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 200
|
||||
config.num_workers = 16
|
||||
config.update_interval = 6
|
||||
config.test_interval = 1
|
||||
config.test_repetitions = 50
|
||||
config.logger = Logger('./log', logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def a3c_cart_pole():
|
||||
config = Config()
|
||||
name = 'CartPole-v0'
|
||||
# name = 'MountainCar-v0'
|
||||
config.task_fn = lambda: ClassicalControl(name, max_steps=200)
|
||||
# config.task_fn = lambda: LunarLander()
|
||||
task = config.task_fn()
|
||||
config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.network_fn = lambda: ActorCriticFCNet(task.state_dim, task.action_dim)
|
||||
config.policy_fn = SamplePolicy
|
||||
config.worker = AdvantageActorCritic
|
||||
config.discount = 0.99
|
||||
config.max_episode_length = 200
|
||||
config.num_workers = 7
|
||||
config.update_interval = 6
|
||||
config.test_interval = 1
|
||||
config.test_repetitions = 30
|
||||
config.logger = Logger('./log', logger)
|
||||
config.gae_tau = 1.0
|
||||
config.entropy_weight = 0.01
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def a2c_cart_pole():
|
||||
config = Config()
|
||||
name = 'CartPole-v0'
|
||||
# name = 'MountainCar-v0'
|
||||
task_fn = lambda: ClassicalControl(name, max_steps=200)
|
||||
# task_fn = lambda: LunarLander()
|
||||
task_fn = lambda **kwargs: ClassicalControl(name, max_steps=200)
|
||||
task = task_fn()
|
||||
config.num_workers = 5
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
||||
@@ -84,8 +41,6 @@ def a2c_cart_pole():
|
||||
config.network_fn = lambda: ActorCriticFCNet(task.state_dim, task.action_dim)
|
||||
config.policy_fn = SamplePolicy
|
||||
config.discount = 0.99
|
||||
config.test_interval = 200
|
||||
config.test_repetitions = 10
|
||||
config.logger = Logger('./log', logger)
|
||||
config.gae_tau = 1.0
|
||||
config.entropy_weight = 0.01
|
||||
@@ -95,14 +50,13 @@ def a2c_cart_pole():
|
||||
def dqn_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 4
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False,
|
||||
history_length=config.history_length)
|
||||
config.task_fn = lambda: PixelAtari(name, frame_skip=4, history_length=config.history_length)
|
||||
action_dim = config.task_fn().action_dim
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
|
||||
config.network_fn = lambda: NatureConvNet(config.history_length, action_dim, gpu=0)
|
||||
# config.network_fn = lambda: DuelingNatureConvNet(config.history_length, action_dim)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8)
|
||||
config.reward_shift_fn = lambda r: np.sign(r)
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 10000
|
||||
@@ -136,63 +90,14 @@ def dqn_ram_atari(name):
|
||||
# config.double_q = False
|
||||
run_episodes(DQNAgent(config))
|
||||
|
||||
def async_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 1
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42)
|
||||
task = config.task_fn()
|
||||
config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001)
|
||||
config.network_fn = lambda: OpenAIConvNet(
|
||||
config.history_length, task.env.action_space.n)
|
||||
config.policy_fn = lambda: StochasticGreedyPolicy(
|
||||
epsilons=[0.7, 0.7, 0.7], final_step=2000000, min_epsilons=[0.1, 0.01, 0.5],
|
||||
probs=[0.4, 0.3, 0.3])
|
||||
# config.worker = OneStepSarsa
|
||||
# config.worker = NStepQLearning
|
||||
config.worker = OneStepQLearning
|
||||
config.reward_shift_fn = lambda r: np.sign(r)
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 10000
|
||||
config.max_episode_length = 10000
|
||||
config.num_workers = 6
|
||||
config.update_interval = 20
|
||||
config.test_interval = 50000
|
||||
config.test_repetitions = 1
|
||||
config.logger = Logger('./log', logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def a3c_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 1
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42)
|
||||
task = config.task_fn()
|
||||
config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001)
|
||||
config.network_fn = lambda: OpenAIActorCriticConvNet(
|
||||
config.history_length, task.env.action_space.n, LSTM=False)
|
||||
config.reward_shift_fn = lambda r: np.sign(r)
|
||||
config.policy_fn = SamplePolicy
|
||||
config.worker = AdvantageActorCritic
|
||||
config.discount = 0.99
|
||||
config.num_workers = 6
|
||||
config.update_interval = 20
|
||||
config.test_interval = 50000
|
||||
config.test_repetitions = 1
|
||||
config.logger = Logger('./log', logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def a2c_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 4
|
||||
config.num_workers = 5
|
||||
task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=84,
|
||||
history_length=config.history_length)
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
||||
task_fn = lambda **kwargs: PixelAtari(name, frame_skip=4, history_length=config.history_length)
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, tag=a2c_pixel_atari.__name__)
|
||||
task = config.task_fn()
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.0007)
|
||||
# config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001)
|
||||
# config.network_fn = lambda: OpenAIActorCriticConvNet(
|
||||
config.network_fn = lambda: NatureActorCriticConvNet(
|
||||
config.history_length, task.task.env.action_space.n, gpu=3)
|
||||
config.reward_shift_fn = lambda r: np.sign(r)
|
||||
@@ -208,98 +113,98 @@ def a2c_pixel_atari(name):
|
||||
config.logger = Logger('./log', logger, skip=True)
|
||||
run_iterations(A2CAgent(config))
|
||||
|
||||
def a3c_continuous():
|
||||
config = Config()
|
||||
config.task_fn = lambda: Pendulum()
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
task = config.task_fn()
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(
|
||||
# lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0),
|
||||
lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True),
|
||||
lambda: GaussianCriticNet(task.state_dim))
|
||||
config.policy_fn = lambda: GaussianPolicy()
|
||||
config.worker = ContinuousAdvantageActorCritic
|
||||
config.discount = 0.99
|
||||
config.num_workers = 8
|
||||
config.update_interval = 20
|
||||
config.test_interval = 1
|
||||
config.test_repetitions = 1
|
||||
config.entropy_weight = 0
|
||||
config.gradient_clip = 40
|
||||
config.logger = Logger('./log', logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def p3o_continuous():
|
||||
config = Config()
|
||||
config.task_fn = lambda: Pendulum()
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolAnt-v1')
|
||||
task = config.task_fn()
|
||||
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim,
|
||||
gpu=-1, unit_std=True)
|
||||
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
|
||||
config.policy_fn = lambda: GaussianPolicy()
|
||||
config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048)
|
||||
config.worker = ProximalPolicyOptimization
|
||||
config.discount = 0.99
|
||||
config.gae_tau = 0.97
|
||||
config.num_workers = 6
|
||||
config.test_interval = 1
|
||||
config.test_repetitions = 1
|
||||
config.entropy_weight = 0
|
||||
config.gradient_clip = 20
|
||||
config.rollout_length = 10000
|
||||
config.optimize_epochs = 1
|
||||
config.ppo_ratio_clip = 0.2
|
||||
config.logger = Logger('./log', logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
|
||||
def d3pg_continuous():
|
||||
config = Config()
|
||||
config.task_fn = lambda: Pendulum()
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolReacher-v1')
|
||||
task = config.task_fn()
|
||||
config.actor_network_fn = lambda: DeterministicActorNet(
|
||||
task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False)
|
||||
config.critic_network_fn = lambda: DeterministicCriticNet(
|
||||
task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False)
|
||||
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4)
|
||||
config.critic_optimizer_fn =\
|
||||
lambda params: torch.optim.Adam(params, lr=1e-4)
|
||||
config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64,
|
||||
state_shape=(task.state_dim, ), action_shape=(task.action_dim, ))
|
||||
config.discount = 0.99
|
||||
config.random_process_fn = \
|
||||
lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2,
|
||||
n_steps_annealing=100000)
|
||||
config.worker = DeterministicPolicyGradient
|
||||
config.num_workers = 6
|
||||
config.min_memory_size = 50
|
||||
config.target_network_mix = 0.001
|
||||
config.test_interval = 500
|
||||
config.test_repetitions = 1
|
||||
config.gradient_clip = 20
|
||||
config.logger = Logger('./log', logger)
|
||||
agent = AsyncAgent(config)
|
||||
agent.run()
|
||||
# def a3c_continuous():
|
||||
# config = Config()
|
||||
# config.task_fn = lambda: Pendulum()
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# task = config.task_fn()
|
||||
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
|
||||
# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
# config.network_fn = lambda: DisjointActorCriticNet(
|
||||
# # lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=False, action_gate=F.tanh, action_scale=2.0),
|
||||
# lambda: GaussianActorNet(task.state_dim, task.action_dim, unit_std=True),
|
||||
# lambda: GaussianCriticNet(task.state_dim))
|
||||
# config.policy_fn = lambda: GaussianPolicy()
|
||||
# config.worker = ContinuousAdvantageActorCritic
|
||||
# config.discount = 0.99
|
||||
# config.num_workers = 8
|
||||
# config.update_interval = 20
|
||||
# config.test_interval = 1
|
||||
# config.test_repetitions = 1
|
||||
# config.entropy_weight = 0
|
||||
# config.gradient_clip = 40
|
||||
# config.logger = Logger('./log', logger)
|
||||
# agent = AsyncAgent(config)
|
||||
# agent.run()
|
||||
#
|
||||
# def p3o_continuous():
|
||||
# config = Config()
|
||||
# config.task_fn = lambda: Pendulum()
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolAnt-v1')
|
||||
# task = config.task_fn()
|
||||
# config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim,
|
||||
# gpu=-1, unit_std=True)
|
||||
# config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim, gpu=-1)
|
||||
# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
# config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
|
||||
#
|
||||
# config.policy_fn = lambda: GaussianPolicy()
|
||||
# config.replay_fn = lambda: GeneralReplay(memory_size=2048, batch_size=2048)
|
||||
# config.worker = ProximalPolicyOptimization
|
||||
# config.discount = 0.99
|
||||
# config.gae_tau = 0.97
|
||||
# config.num_workers = 6
|
||||
# config.test_interval = 1
|
||||
# config.test_repetitions = 1
|
||||
# config.entropy_weight = 0
|
||||
# config.gradient_clip = 20
|
||||
# config.rollout_length = 10000
|
||||
# config.optimize_epochs = 1
|
||||
# config.ppo_ratio_clip = 0.2
|
||||
# config.logger = Logger('./log', logger)
|
||||
# agent = AsyncAgent(config)
|
||||
# agent.run()
|
||||
#
|
||||
# def d3pg_continuous():
|
||||
# config = Config()
|
||||
# config.task_fn = lambda: Pendulum()
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalker-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('BipedalWalkerHardcore-v2')
|
||||
# # config.task_fn = lambda: Box2DContinuous('LunarLanderContinuous-v2')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1')
|
||||
# # config.task_fn = lambda: Roboschool('RoboschoolReacher-v1')
|
||||
# task = config.task_fn()
|
||||
# config.actor_network_fn = lambda: DeterministicActorNet(
|
||||
# task.state_dim, task.action_dim, F.tanh, 2, non_linear=F.relu, batch_norm=False)
|
||||
# config.critic_network_fn = lambda: DeterministicCriticNet(
|
||||
# task.state_dim, task.action_dim, non_linear=F.relu, batch_norm=False)
|
||||
# config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
|
||||
# config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, lr=1e-4)
|
||||
# config.critic_optimizer_fn =\
|
||||
# lambda params: torch.optim.Adam(params, lr=1e-4)
|
||||
# config.replay_fn = lambda: SharedReplay(memory_size=1000000, batch_size=64,
|
||||
# state_shape=(task.state_dim, ), action_shape=(task.action_dim, ))
|
||||
# config.discount = 0.99
|
||||
# config.random_process_fn = \
|
||||
# lambda: OrnsteinUhlenbeckProcess(size=task.action_dim, theta=0.15, sigma=0.2,
|
||||
# n_steps_annealing=100000)
|
||||
# config.worker = DeterministicPolicyGradient
|
||||
# config.num_workers = 6
|
||||
# config.min_memory_size = 50
|
||||
# config.target_network_mix = 0.001
|
||||
# config.test_interval = 500
|
||||
# config.test_repetitions = 1
|
||||
# config.gradient_clip = 20
|
||||
# config.logger = Logger('./log', logger)
|
||||
# agent = AsyncAgent(config)
|
||||
# agent.run()
|
||||
|
||||
def ddpg_continuous():
|
||||
config = Config()
|
||||
@@ -359,8 +264,7 @@ def categorical_dqn_cart_pole():
|
||||
def categorical_dqn_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 4
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False,
|
||||
history_length=config.history_length)
|
||||
config.task_fn = lambda: PixelAtari(name, frame_skip=4, history_length=config.history_length)
|
||||
action_dim = config.task_fn().action_dim
|
||||
config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.00025, eps=0.01 / 32)
|
||||
config.network_fn = lambda: CategoricalConvNet(config.history_length, action_dim, config.categorical_n_atoms, gpu=0)
|
||||
@@ -381,7 +285,7 @@ def categorical_dqn_pixel_atari(name):
|
||||
|
||||
def n_step_dqn_cart_pole():
|
||||
config = Config()
|
||||
task_fn = lambda: ClassicalControl('CartPole-v0', max_steps=200)
|
||||
task_fn = lambda **kwargs: ClassicalControl('CartPole-v0', max_steps=200)
|
||||
task = task_fn()
|
||||
config.num_workers = 5
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
||||
@@ -397,11 +301,10 @@ def n_step_dqn_cart_pole():
|
||||
def n_step_dqn_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 4
|
||||
task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=True,
|
||||
history_length=config.history_length)
|
||||
task_fn = lambda **kwargs: PixelAtari(name, frame_skip=4, history_length=config.history_length)
|
||||
task = task_fn()
|
||||
config.num_workers = 8
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers, tag=n_step_dqn_pixel_atari.__name__)
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
|
||||
config.network_fn = lambda: NatureConvNet(config.history_length, task.action_dim, gpu=0)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
|
||||
@@ -433,8 +336,7 @@ def quantile_regression_dqn_cart_pole():
|
||||
def quantile_regression_dqn_pixel_atari(name):
|
||||
config = Config()
|
||||
config.history_length = 4
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False,
|
||||
history_length=config.history_length)
|
||||
config.task_fn = lambda: PixelAtari(name, frame_skip=4, history_length=config.history_length)
|
||||
action_dim = config.task_fn().action_dim
|
||||
config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.00005, eps=0.01 / 32)
|
||||
config.network_fn = lambda: QuantileConvNet(config.history_length, action_dim, config.num_quantiles, gpu=0)
|
||||
@@ -460,28 +362,19 @@ if __name__ == '__main__':
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
# dqn_cart_pole()
|
||||
# a2c_cart_pole()
|
||||
# categorical_dqn_cart_pole()
|
||||
# quantile_regression_dqn_cart_pole()
|
||||
# async_cart_pole()
|
||||
# a3c_cart_pole()
|
||||
a2c_cart_pole()
|
||||
# a3c_continuous()
|
||||
# p3o_continuous()
|
||||
# d3pg_continuous()
|
||||
# ddpg_continuous()
|
||||
# n_step_dqn_cart_pole()
|
||||
|
||||
# dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
# categorical_dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
# quantile_regression_dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
# n_step_dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
# async_pixel_atari('PongNoFrameskip-v4')
|
||||
# a3c_pixel_atari('PongNoFrameskip-v4')
|
||||
# a2c_pixel_atari('PongNoFrameskip-v4')
|
||||
# categorical_dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
quantile_regression_dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
# n_step_dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
|
||||
# dqn_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# async_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# a3c_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
|
||||
# dqn_ram_atari('Pong-ramNoFrameskip-v4')
|
||||
|
||||
|
||||
@@ -106,3 +106,78 @@ class QuantileNet(BasicNet):
|
||||
def predict(self, x, to_numpy=False):
|
||||
quantiles = self.forward(x)
|
||||
return quantiles.view((-1, self.n_actions, self.n_quantiles))
|
||||
|
||||
class GammaNet(BasicNet):
|
||||
def predict(self, features, aux_features):
|
||||
attention = self.compute_attention(features)
|
||||
|
||||
aux_features = torch.stack(aux_features)
|
||||
aux_features = aux_features * attention.t().unsqueeze(-1)
|
||||
aux_features = aux_features.transpose(0, 1).contiguous().sum(1)
|
||||
|
||||
phi = features + aux_features
|
||||
pre_prob = self.fc_actor(phi)
|
||||
prob = F.softmax(pre_prob, dim=1)
|
||||
log_prob = F.log_softmax(pre_prob, dim=1)
|
||||
value = self.fc_critic(phi)
|
||||
return prob, log_prob, value
|
||||
|
||||
def compute_attention(self, phi):
|
||||
attention = self.fc_attention(phi)
|
||||
attention = F.sigmoid(attention)
|
||||
return attention
|
||||
|
||||
def q(self, x):
|
||||
return self.fc_q(x)
|
||||
|
||||
def predict(self, features, aux_features):
|
||||
aux_features.append(features)
|
||||
phi = torch.cat(aux_features, dim=1)
|
||||
|
||||
pre_prob = self.fc_actor(phi)
|
||||
prob = F.softmax(pre_prob, dim=1)
|
||||
log_prob = F.log_softmax(pre_prob, dim=1)
|
||||
value = self.fc_critic(phi)
|
||||
return prob, log_prob, value
|
||||
|
||||
def feature(self, x):
|
||||
return self.forward(x)
|
||||
|
||||
|
||||
class GammaAttentionNet(BasicNet):
|
||||
def predict(self, features, aux_features):
|
||||
attention = self.compute_attention(features)
|
||||
|
||||
aux_features = torch.stack(aux_features)
|
||||
aux_features = aux_features * attention.t().unsqueeze(-1)
|
||||
aux_features = aux_features.transpose(0, 1).contiguous().sum(1)
|
||||
|
||||
phi = features + aux_features
|
||||
pre_prob = self.fc_actor(phi)
|
||||
prob = F.softmax(pre_prob, dim=1)
|
||||
log_prob = F.log_softmax(pre_prob, dim=1)
|
||||
value = self.fc_critic(phi)
|
||||
return prob, log_prob, value
|
||||
|
||||
def compute_attention(self, phi):
|
||||
attention = self.fc_attention(phi)
|
||||
# attention = F.relu(attention)
|
||||
# attention = F.tanh(attention)
|
||||
# attention = (attention + 1) / 0.5
|
||||
# attention = F.tanh(attention)
|
||||
# attention = F.tanh(attention)
|
||||
# attention = F.sigmoid(attention)
|
||||
attention = F.softmax(attention, dim=1)
|
||||
# max_attention = 10
|
||||
# cond = (attention < max_attention).float().detach()
|
||||
# attention = attention * cond + max_attention * (1 - cond)
|
||||
# cond = (attention > -max_attention).float().detach()
|
||||
# attention = attention * cond + -max_attention * (1 - cond)
|
||||
# self.attention = attention.data.cpu().numpy()
|
||||
return attention
|
||||
|
||||
def q(self, x):
|
||||
return self.fc_q(x)
|
||||
|
||||
def feature(self, x):
|
||||
return self.forward(x)
|
||||
|
||||
+57
-1
@@ -183,4 +183,60 @@ class QuantileConvNet(nn.Module, QuantileNet):
|
||||
y = y.view(y.size(0), -1)
|
||||
y = F.relu(self.fc4(y))
|
||||
y = self.fc5(y)
|
||||
return y
|
||||
return y
|
||||
|
||||
class GammaConvNet(nn.Module, GammaNet):
|
||||
def __init__(self, in_channels, action_dim, num_peers, gpu=-1):
|
||||
super(GammaConvNet, self).__init__()
|
||||
hidden_size = 512
|
||||
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4)
|
||||
self.conv2 = nn.Conv2d(32, 64, kernel_size=4, stride=2)
|
||||
self.conv3 = nn.Conv2d(64, 32, kernel_size=3, stride=1)
|
||||
self.fc4 = nn.Linear(7 * 7 * 32, hidden_size)
|
||||
|
||||
self.fc_actor = nn.Linear(hidden_size * num_peers, action_dim)
|
||||
self.fc_critic = nn.Linear(hidden_size * num_peers, 1)
|
||||
|
||||
self.fc_attention = nn.Linear(hidden_size, num_peers - 1)
|
||||
self.fc_q = nn.Linear(hidden_size, action_dim)
|
||||
|
||||
self.fc_actor_main = nn.Linear(hidden_size, action_dim)
|
||||
self.fc_critic_main = nn.Linear(hidden_size, 1)
|
||||
self.compute_attention = self.softmax_attention
|
||||
BasicNet.__init__(self, gpu=gpu)
|
||||
|
||||
def forward(self, x, update_lstm=True):
|
||||
x = self.variable(x)
|
||||
x = F.relu(self.conv1(x))
|
||||
x = F.relu(self.conv2(x))
|
||||
x = F.relu(self.conv3(x))
|
||||
x = x.view(x.size(0), -1)
|
||||
phi = F.relu(self.fc4(x))
|
||||
return phi
|
||||
|
||||
class GammaAttentionConvNet(nn.Module, GammaAttentionNet):
|
||||
def __init__(self, in_channels, action_dim, num_peers, gpu=-1):
|
||||
super(GammaAttentionConvNet, self).__init__()
|
||||
hidden_size = 512
|
||||
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4)
|
||||
self.conv2 = nn.Conv2d(32, 64, kernel_size=4, stride=2)
|
||||
self.conv3 = nn.Conv2d(64, 32, kernel_size=3, stride=1)
|
||||
self.fc4 = nn.Linear(7 * 7 * 32, hidden_size)
|
||||
|
||||
self.fc_actor = nn.Linear(hidden_size, action_dim)
|
||||
self.fc_critic = nn.Linear(hidden_size, 1)
|
||||
|
||||
self.fc_attention = nn.Linear(hidden_size, num_peers - 1)
|
||||
self.fc_q = nn.Linear(hidden_size, action_dim)
|
||||
|
||||
BasicNet.__init__(self, gpu=gpu)
|
||||
self.fc_attention.weight.data.zero_()
|
||||
|
||||
def forward(self, x, update_lstm=True):
|
||||
x = self.variable(x)
|
||||
x = F.relu(self.conv1(x))
|
||||
x = F.relu(self.conv2(x))
|
||||
x = F.relu(self.conv3(x))
|
||||
x = x.view(x.size(0), -1)
|
||||
phi = F.relu(self.fc4(x))
|
||||
return phi
|
||||
|
||||
@@ -40,7 +40,7 @@ class DuelingFCNet(nn.Module, DuelingNet):
|
||||
|
||||
# Network for CartPole with actor critic
|
||||
class ActorCriticFCNet(nn.Module, ActorCriticNet):
|
||||
def __init__(self, state_dim, action_dim):
|
||||
def __init__(self, state_dim, action_dim, gpu=-1):
|
||||
super(ActorCriticFCNet, self).__init__()
|
||||
hidden_size1 = 64
|
||||
hidden_size2 = 64
|
||||
@@ -48,7 +48,7 @@ class ActorCriticFCNet(nn.Module, ActorCriticNet):
|
||||
self.fc2 = nn.Linear(hidden_size1, hidden_size2)
|
||||
self.fc_actor = nn.Linear(hidden_size2, action_dim)
|
||||
self.fc_critic = nn.Linear(hidden_size2, 1)
|
||||
BasicNet.__init__(self, False)
|
||||
BasicNet.__init__(self, gpu=gpu)
|
||||
|
||||
def forward(self, x, update_LSTM=True):
|
||||
x = self.variable(x)
|
||||
@@ -91,4 +91,47 @@ class QuantileFCNet(nn.Module, QuantileNet):
|
||||
phi = F.relu(self.fc1(x))
|
||||
phi = F.relu(self.fc2(phi))
|
||||
quantiles = self.fc3(phi)
|
||||
return quantiles
|
||||
return quantiles
|
||||
|
||||
class GammaFCNet(nn.Module, GammaNet):
|
||||
def __init__(self, state_dim, action_dim, num_peers, gpu=-1):
|
||||
super(GammaFCNet, self).__init__()
|
||||
hidden_size = 64
|
||||
self.fc1 = nn.Linear(state_dim, hidden_size)
|
||||
self.fc2 = nn.Linear(hidden_size, hidden_size)
|
||||
self.fc_actor = nn.Linear(hidden_size * num_peers, action_dim)
|
||||
self.fc_critic = nn.Linear(hidden_size * num_peers, 1)
|
||||
|
||||
self.fc_attention = nn.Linear(hidden_size, num_peers - 1)
|
||||
self.fc_q = nn.Linear(hidden_size, action_dim)
|
||||
|
||||
# self.fc_actor_main = nn.Linear(hidden_size, action_dim)
|
||||
# self.fc_critic_main = nn.Linear(hidden_size, 1)
|
||||
self.compute_attention = self.softmax_attention
|
||||
BasicNet.__init__(self, gpu=gpu)
|
||||
|
||||
def forward(self, x, update_lstm=True):
|
||||
x = self.variable(x)
|
||||
x = F.relu(self.fc1(x))
|
||||
x = F.relu(self.fc2(x))
|
||||
return x
|
||||
|
||||
class GammaAttentionFCNet(nn.Module, GammaAttentionNet):
|
||||
def __init__(self, state_dim, action_dim, num_peers, gpu=-1):
|
||||
super(GammaAttentionFCNet, self).__init__()
|
||||
hidden_size = 64
|
||||
self.fc1 = nn.Linear(state_dim, hidden_size)
|
||||
self.fc2 = nn.Linear(hidden_size, hidden_size)
|
||||
self.fc_actor = nn.Linear(hidden_size, action_dim)
|
||||
self.fc_critic = nn.Linear(hidden_size, 1)
|
||||
|
||||
self.fc_attention = nn.Linear(hidden_size, num_peers - 1)
|
||||
self.fc_q = nn.Linear(hidden_size, action_dim)
|
||||
BasicNet.__init__(self, gpu=gpu)
|
||||
self.fc_attention.weight.data.zero_()
|
||||
|
||||
def forward(self, x, update_lstm=True):
|
||||
x = self.variable(x)
|
||||
x = F.relu(self.fc1(x))
|
||||
x = F.relu(self.fc2(x))
|
||||
return x
|
||||
|
||||
+1
-1
@@ -24,7 +24,7 @@ class Config:
|
||||
self.exploration_steps = 0
|
||||
self.logger = None
|
||||
self.history_length = 1
|
||||
self.test_interval = 100
|
||||
self.test_interval = 0
|
||||
self.test_repetitions = 50
|
||||
self.double_q = False
|
||||
self.tag = 'vanilla'
|
||||
|
||||
@@ -77,7 +77,17 @@ def run_iterations(agent):
|
||||
pickle.dump({'rewards': rewards,
|
||||
'steps': steps}, f)
|
||||
agent.save('data/%s-%s-model-%s.bin' % (agent_name, config.tag, agent.task.name))
|
||||
if config.test_interval and iteration % config.test_interval == 0:
|
||||
test_rewards, test_steps = agent.evaluate()
|
||||
config.logger.info('total steps %d, test reward %f, test steps %d' % (
|
||||
agent.total_steps, test_rewards, test_steps
|
||||
))
|
||||
iteration += 1
|
||||
if config.max_steps and agent.total_steps >= config.max_steps:
|
||||
agent.close()
|
||||
break
|
||||
|
||||
return steps, rewards
|
||||
|
||||
def sync_grad(target_network, src_network):
|
||||
for param, src_param in zip(target_network.parameters(), src_network.parameters()):
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ def plot_curves(xy_list, xaxis, title):
|
||||
plt.ylabel("Episode Rewards")
|
||||
plt.tight_layout()
|
||||
|
||||
def plot_results(dirs, num_timesteps, xaxis, task_name):
|
||||
def plot_results(dirs, num_timesteps=1e8, xaxis=X_TIMESTEPS, task_name=''):
|
||||
tslist = []
|
||||
for dir in dirs:
|
||||
ts = load_results(dir)
|
||||
|
||||
Reference in New Issue
Block a user