mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Refactor replay
This commit is contained in:
+34
-230
@@ -10,260 +10,64 @@ import random
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
class Replay:
|
||||
def __init__(self, memory_size, batch_size, dtype=np.float32):
|
||||
self.memory_size = memory_size
|
||||
self.batch_size = batch_size
|
||||
self.dtype = dtype
|
||||
|
||||
self.states = None
|
||||
self.actions = np.empty(self.memory_size, dtype=np.uint8)
|
||||
self.rewards = np.empty(self.memory_size)
|
||||
self.next_states = None
|
||||
self.terminals = np.empty(self.memory_size, dtype=np.uint8)
|
||||
|
||||
self.pos = 0
|
||||
self.full = False
|
||||
|
||||
|
||||
def feed(self, experience):
|
||||
state, action, reward, next_state, done = experience
|
||||
|
||||
if self.states is None:
|
||||
self.states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype)
|
||||
self.next_states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype)
|
||||
|
||||
self.states[self.pos][:] = state
|
||||
self.actions[self.pos] = action
|
||||
self.rewards[self.pos] = reward
|
||||
self.next_states[self.pos][:] = next_state
|
||||
self.terminals[self.pos] = done
|
||||
|
||||
self.pos += 1
|
||||
if self.pos == self.memory_size:
|
||||
self.full = True
|
||||
self.pos = 0
|
||||
|
||||
def sample(self):
|
||||
upper_bound = self.memory_size if self.full else self.pos
|
||||
sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size)
|
||||
return [self.states[sampled_indices],
|
||||
self.actions[sampled_indices],
|
||||
self.rewards[sampled_indices],
|
||||
self.next_states[sampled_indices],
|
||||
self.terminals[sampled_indices]]
|
||||
|
||||
class HybridRewardReplay:
|
||||
def __init__(self, memory_size, batch_size, dtype=np.float32):
|
||||
self.memory_size = memory_size
|
||||
self.batch_size = batch_size
|
||||
self.dtype = dtype
|
||||
|
||||
self.states = None
|
||||
self.actions = np.empty(self.memory_size, dtype=np.uint8)
|
||||
self.rewards = None
|
||||
self.next_states = None
|
||||
self.terminals = np.empty(self.memory_size, dtype=np.uint8)
|
||||
|
||||
self.pos = 0
|
||||
self.full = False
|
||||
|
||||
|
||||
def feed(self, experience):
|
||||
state, action, reward, next_state, done = experience
|
||||
|
||||
if self.states is None:
|
||||
self.rewards = np.empty((self.memory_size, ) + reward.shape, dtype=self.dtype)
|
||||
self.states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype)
|
||||
self.next_states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype)
|
||||
|
||||
self.states[self.pos][:] = state
|
||||
self.actions[self.pos] = action
|
||||
self.rewards[self.pos][:] = reward
|
||||
self.next_states[self.pos][:] = next_state
|
||||
self.terminals[self.pos] = done
|
||||
|
||||
self.pos += 1
|
||||
if self.pos == self.memory_size:
|
||||
self.full = True
|
||||
self.pos = 0
|
||||
|
||||
def sample(self):
|
||||
upper_bound = self.memory_size if self.full else self.pos
|
||||
sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size)
|
||||
return [self.states[sampled_indices],
|
||||
self.actions[sampled_indices],
|
||||
self.rewards[sampled_indices],
|
||||
self.next_states[sampled_indices],
|
||||
self.terminals[sampled_indices]]
|
||||
|
||||
class SharedReplay:
|
||||
def __init__(self, memory_size, batch_size, state_shape, action_shape):
|
||||
self.memory_size = memory_size
|
||||
self.batch_size = batch_size
|
||||
|
||||
self.states = torch.zeros((self.memory_size, ) + state_shape)
|
||||
self.actions = torch.zeros((self.memory_size, ) + action_shape)
|
||||
self.rewards = torch.zeros(self.memory_size)
|
||||
self.next_states = torch.zeros((self.memory_size, ) + state_shape)
|
||||
self.terminals = torch.zeros(self.memory_size)
|
||||
|
||||
self.states.share_memory_()
|
||||
self.actions.share_memory_()
|
||||
self.rewards.share_memory_()
|
||||
self.next_states.share_memory_()
|
||||
self.terminals.share_memory_()
|
||||
|
||||
self.pos = 0
|
||||
self.full = False
|
||||
self.buffer_lock = mp.Lock()
|
||||
|
||||
def feed_(self, experience):
|
||||
state, action, reward, next_state, done = experience
|
||||
self.states[self.pos][:] = torch.FloatTensor(state)
|
||||
self.actions[self.pos][:] = torch.FloatTensor(action)
|
||||
self.rewards[self.pos] = reward
|
||||
self.next_states[self.pos][:] = torch.FloatTensor(next_state)
|
||||
self.terminals[self.pos] = done
|
||||
|
||||
self.pos += 1
|
||||
if self.pos == self.memory_size:
|
||||
self.full = True
|
||||
self.pos = 0
|
||||
|
||||
def size(self):
|
||||
if self.full:
|
||||
return self.memory_size
|
||||
return self.pos
|
||||
|
||||
def sample_(self):
|
||||
upper_bound = self.memory_size if self.full else self.pos
|
||||
sampled_indices = torch.LongTensor(np.random.randint(0, upper_bound, size=self.batch_size))
|
||||
return [self.states[sampled_indices],
|
||||
self.actions[sampled_indices],
|
||||
self.rewards[sampled_indices],
|
||||
self.next_states[sampled_indices],
|
||||
self.terminals[sampled_indices]]
|
||||
|
||||
def feed(self, experience):
|
||||
with self.buffer_lock:
|
||||
self.feed_(experience)
|
||||
|
||||
def sample(self):
|
||||
with self.buffer_lock:
|
||||
return self.sample_()
|
||||
|
||||
def state_dict(self):
|
||||
return dict((key, getattr(self, key)) for key in ['actions', 'states', 'rewards', 'next_states', 'terminals', 'pos'])
|
||||
|
||||
def load_state_dict(self, state):
|
||||
for key in ['actions', 'states', 'rewards', 'next_states', 'terminals', 'pos']:
|
||||
val = state[key]
|
||||
setattr(self, key, val)
|
||||
|
||||
def save(self, file_name):
|
||||
with open(file_name, 'wb') as f:
|
||||
torch.save(self.state_dict(), f)
|
||||
|
||||
def load(self, file_name):
|
||||
state = torch.load(file_name)
|
||||
self.load_state_dict(state)
|
||||
|
||||
class HighDimActionReplay:
|
||||
def __init__(self, memory_size, batch_size, dtype=np.float32):
|
||||
self.memory_size = memory_size
|
||||
self.batch_size = batch_size
|
||||
self.dtype = dtype
|
||||
|
||||
self.states = None
|
||||
self.actions = None
|
||||
self.rewards = np.empty(self.memory_size)
|
||||
self.next_states = None
|
||||
self.terminals = np.empty(self.memory_size, dtype=np.int8)
|
||||
|
||||
self.pos = 0
|
||||
self.full = False
|
||||
|
||||
|
||||
def feed(self, experience):
|
||||
state, action, reward, next_state, done = experience
|
||||
|
||||
if self.states is None:
|
||||
self.states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype)
|
||||
self.actions = np.empty((self.memory_size, ) + action.shape)
|
||||
self.next_states = np.empty((self.memory_size, ) + state.shape, dtype=self.dtype)
|
||||
|
||||
self.states[self.pos][:] = state
|
||||
self.actions[self.pos][:] = action
|
||||
self.rewards[self.pos] = reward
|
||||
self.next_states[self.pos][:] = next_state
|
||||
self.terminals[self.pos] = done
|
||||
|
||||
self.pos += 1
|
||||
if self.pos == self.memory_size:
|
||||
self.full = True
|
||||
self.pos = 0
|
||||
|
||||
def size(self):
|
||||
if self.full:
|
||||
return self.memory_size
|
||||
return self.pos
|
||||
|
||||
def sample(self):
|
||||
upper_bound = self.memory_size if self.full else self.pos
|
||||
sampled_indices = np.random.randint(0, upper_bound, size=self.batch_size)
|
||||
return [self.states[sampled_indices],
|
||||
self.actions[sampled_indices],
|
||||
self.rewards[sampled_indices],
|
||||
self.next_states[sampled_indices],
|
||||
self.terminals[sampled_indices]]
|
||||
|
||||
class GeneralReplay:
|
||||
def __init__(self, memory_size, batch_size):
|
||||
self.buffer = []
|
||||
self.memory_size = memory_size
|
||||
self.batch_size = batch_size
|
||||
self.data = None
|
||||
|
||||
def feed(self, experiences):
|
||||
for experience in zip(*experiences):
|
||||
self.feed_single(experience)
|
||||
self.pos = 0
|
||||
self.full = False
|
||||
|
||||
def feed_single(self, experience):
|
||||
self.buffer.append(experience)
|
||||
if len(self.buffer) > self.memory_size:
|
||||
del self.buffer[0]
|
||||
def feed(self, experience):
|
||||
if self.data is None:
|
||||
self.data = []
|
||||
for unit in experience:
|
||||
if np.isscalar(unit):
|
||||
self.data.append(np.zeros(self.memory_size, dtype=type(unit)))
|
||||
else:
|
||||
self.data.append(np.zeros((self.memory_size, ) + unit.shape, unit.dtype))
|
||||
for buffer_unit, exp_unit in zip(self.data, experience):
|
||||
buffer_unit[self.pos] = exp_unit
|
||||
|
||||
self.pos += 1
|
||||
if self.pos == self.memory_size:
|
||||
self.full = True
|
||||
self.pos = 0
|
||||
|
||||
def feed_batch(self, experience):
|
||||
experience = zip(*experience)
|
||||
for exp in experience:
|
||||
self.feed(exp)
|
||||
|
||||
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):
|
||||
self.buffer = []
|
||||
|
||||
def full(self):
|
||||
return len(self.buffer) == self.memory_size
|
||||
upper_bound = self.memory_size if self.full else self.pos
|
||||
sampled_indices = np.random.randint(0, upper_bound, size=batch_size)
|
||||
return [unit[sampled_indices] for unit in self.data]
|
||||
|
||||
def size(self):
|
||||
return len(self.buffer)
|
||||
if self.full:
|
||||
return self.memory_size
|
||||
return self.pos
|
||||
|
||||
def empty(self):
|
||||
return not len(self.buffer)
|
||||
return not self.full and not self.pos
|
||||
|
||||
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.non_zero_reward = Replay(memory_size, batch_size / 2)
|
||||
self.zero_reward = Replay(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)
|
||||
self.zero_reward.feed(exp)
|
||||
else:
|
||||
self.non_zero_reward.feed_single(exp)
|
||||
self.non_zero_reward.feed(exp)
|
||||
|
||||
def sample(self):
|
||||
if self.zero_reward.empty():
|
||||
|
||||
@@ -132,7 +132,7 @@ def dqn_pixel_atari(name):
|
||||
config.network_fn = lambda state_dim, action_dim: VanillaNet(action_dim, NatureConvBody(), gpu=0)
|
||||
# config.network_fn = lambda state_dim, action_dim: DuelingNet(action_dim, NatureConvBody(), gpu=0)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32)
|
||||
config.state_normalizer = ImageNormalizer()
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
config.discount = 0.99
|
||||
@@ -173,7 +173,7 @@ def categorical_dqn_pixel_atari(name):
|
||||
config.network_fn = lambda state_dim, action_dim: \
|
||||
CategoricalNet(action_dim, config.categorical_n_atoms, NatureConvBody(), gpu=1)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32)
|
||||
config.discount = 0.99
|
||||
config.state_normalizer = ImageNormalizer()
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
@@ -195,7 +195,7 @@ def quantile_regression_dqn_pixel_atari(name):
|
||||
config.network_fn = lambda state_dim, action_dim: \
|
||||
QuantileNet(action_dim, config.num_quantiles, NatureConvBody(), gpu=2)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.01)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32)
|
||||
config.state_normalizer = ImageNormalizer()
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
config.discount = 0.99
|
||||
@@ -258,7 +258,7 @@ def dqn_ram_atari(name):
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
|
||||
config.network_fn = lambda state_dim, action_dim: VanillaNet(action_dim, TwoLayerFCBody(state_dim), gpu=2)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=0.1, final_step=1000000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32, dtype=np.uint8)
|
||||
config.replay_fn = lambda: Replay(memory_size=100000, batch_size=32)
|
||||
config.state_normalizer = RescaleNormalizer(1.0 / 128)
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
config.discount = 0.99
|
||||
@@ -310,17 +310,17 @@ def ddpg_continuous():
|
||||
# config.task_fn = lambda: Pendulum(log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolInvertedPendulum-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolReacher-v1', log_dir=log_dir)
|
||||
config.task_fn = lambda: Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
|
||||
config.task_fn = lambda: Roboschool('RoboschoolHopper-v1')
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolAnt-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: Roboschool('RoboschoolWalker2d-v1', log_dir=log_dir)
|
||||
# config.task_fn = lambda: DMControl('cartpole', 'balance', log_dir=log_dir)
|
||||
# config.task_fn = lambda: DMControl('finger', 'spin', log_dir=log_dir)
|
||||
config.evaluation_env = config.task_fn()
|
||||
config.evaluation_env = Roboschool('RoboschoolHopper-v1', log_dir=log_dir)
|
||||
config.actor_network_fn = lambda state_dim, action_dim: DeterministicActorNet(state_dim, action_dim)
|
||||
config.critic_network_fn = lambda state_dim, action_dim: DeterministicCriticNet(state_dim, action_dim)
|
||||
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: HighDimActionReplay(memory_size=1000000, batch_size=64)
|
||||
config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=64)
|
||||
config.discount = 0.99
|
||||
config.state_normalizer = RunningStatsNormalizer()
|
||||
config.random_process_fn = \
|
||||
|
||||
Reference in New Issue
Block a user