mirror of
https://github.com/wassname/kair_algorithms_draft.git
synced 2026-10-08 12:08:53 +08:00
Add random initial action in ddpg (#13)
* Add random initial actions in ddpg * Add reacher-v2 example of ddpg
This commit is contained in:
1 parent
ecb42d30d2
commit
d2b670015c
9 files changed
+272
-40
No files matched your search
@@ -75,24 +75,24 @@ class Agent(AbstractAgent):
|
||||
self.load_params(args.load_from)
|
||||
|
||||
# replay memory
|
||||
self.beta = self.hyper_params["PER_BETA"]
|
||||
self.memory = PrioritizedReplayBuffer(
|
||||
self.hyper_params["BUFFER_SIZE"],
|
||||
self.hyper_params["BATCH_SIZE"],
|
||||
alpha=self.hyper_params["PER_ALPHA"],
|
||||
)
|
||||
if not self.args.test:
|
||||
self.beta = self.hyper_params["PER_BETA"]
|
||||
self.memory = PrioritizedReplayBuffer(
|
||||
self.hyper_params["BUFFER_SIZE"],
|
||||
self.hyper_params["BATCH_SIZE"],
|
||||
alpha=self.hyper_params["PER_ALPHA"],
|
||||
)
|
||||
|
||||
def select_action(self, state: np.ndarray) -> torch.Tensor:
|
||||
"""Select an action from the input space."""
|
||||
self.curr_state = state
|
||||
|
||||
state = torch.FloatTensor(state).to(device)
|
||||
selected_action = self.actor(state)
|
||||
|
||||
if not self.args.test:
|
||||
selected_action = self.actor(state)
|
||||
selected_action += torch.FloatTensor(self.noise.sample()).to(device)
|
||||
|
||||
selected_action = torch.clamp(selected_action, -1.0, 1.0)
|
||||
selected_action = torch.clamp(selected_action, -1.0, 1.0)
|
||||
|
||||
return selected_action
|
||||
|
||||
@@ -101,7 +101,8 @@ class Agent(AbstractAgent):
|
||||
action = action.detach().cpu().numpy()
|
||||
next_state, reward, done, _ = self.env.step(action)
|
||||
|
||||
self.memory.add(self.curr_state, action, reward, next_state, done)
|
||||
if not self.args.test:
|
||||
self.memory.add(self.curr_state, action, reward, next_state, done)
|
||||
|
||||
return next_state, reward, done
|
||||
|
||||
|
||||
Reference in new issue
Block a user