mirror of
https://github.com/wassname/kair_algorithms_draft.git
synced 2026-09-09 11:25:10 +08:00
Convert code to python 2.7 (#35)
* Convert code format to python2.7 (SAC) * Convert code format python2.7 (TD3, all fD) * Remove no use import and black setting * Change SAC param * Change env name Reacher-v2 to v1 * Remove old version reacher training script * Convert code format python2.7 * Modify .travis.yml * Add install command python3.6 & black on Makefile * Fix seperator to tab on Makefile * Modify Makefile * Fix little error * Change td3 gamma parameter
This commit is contained in:
@@ -10,7 +10,6 @@
|
||||
"""
|
||||
|
||||
import pickle
|
||||
from typing import List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -39,6 +38,8 @@ class Agent(SACAgent):
|
||||
|
||||
if not self.args.test:
|
||||
# load demo replay memory
|
||||
# TODO: should make new demo to set protocol 2
|
||||
# e.g. pickle.dump(your_object, your_file, protocol=2)
|
||||
with open(self.args.demo_path, "rb") as f:
|
||||
demos = pickle.load(f)
|
||||
|
||||
@@ -65,7 +66,7 @@ class Agent(SACAgent):
|
||||
epsilon_d=self.hyper_params["PER_EPS_DEMO"],
|
||||
)
|
||||
|
||||
def _add_transition_to_memory(self, transition: Tuple[np.ndarray, ...]):
|
||||
def _add_transition_to_memory(self, transition):
|
||||
"""Add 1 step and n step transitions to memory."""
|
||||
# add n-step transition
|
||||
if self.use_n_step:
|
||||
@@ -77,19 +78,7 @@ class Agent(SACAgent):
|
||||
self.memory.add(*transition)
|
||||
|
||||
# pylint: disable=too-many-statements
|
||||
def update_model(
|
||||
self,
|
||||
experiences: Tuple[
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
List[int],
|
||||
],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
def update_model(self, experiences):
|
||||
"""Train the model after each episode."""
|
||||
states, actions, rewards, next_states, dones, weights, indices, eps_d = (
|
||||
experiences
|
||||
@@ -212,7 +201,7 @@ class Agent(SACAgent):
|
||||
def pretrain(self):
|
||||
"""Pretraining steps."""
|
||||
pretrain_loss = list()
|
||||
print("[INFO] Pre-Train %d steps." % self.hyper_params["PRETRAIN_STEP"])
|
||||
print ("[INFO] Pre-Train %d steps." % self.hyper_params["PRETRAIN_STEP"])
|
||||
for i_step in range(1, self.hyper_params["PRETRAIN_STEP"] + 1):
|
||||
loss = self.update_model()
|
||||
pretrain_loss.append(loss) # for logging
|
||||
|
||||
@@ -9,7 +9,6 @@
|
||||
"""
|
||||
|
||||
import pickle
|
||||
from typing import List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -38,6 +37,8 @@ class Agent(TD3Agent):
|
||||
|
||||
if not self.args.test:
|
||||
# load demo replay memory
|
||||
# TODO: should make new demo to set protocol 2
|
||||
# e.g. pickle.dump(your_object, your_file, protocol=2)
|
||||
with open(self.args.demo_path, "rb") as f:
|
||||
demos = pickle.load(f)
|
||||
|
||||
@@ -64,7 +65,7 @@ class Agent(TD3Agent):
|
||||
epsilon_d=self.hyper_params["PER_EPS_DEMO"],
|
||||
)
|
||||
|
||||
def _add_transition_to_memory(self, transition: Tuple[np.ndarray, ...]):
|
||||
def _add_transition_to_memory(self, transition):
|
||||
"""Add 1 step and n step transitions to memory."""
|
||||
# add n-step transition
|
||||
if self.use_n_step:
|
||||
@@ -75,9 +76,7 @@ class Agent(TD3Agent):
|
||||
if transition:
|
||||
self.memory.add(*transition)
|
||||
|
||||
def _get_critic_loss(
|
||||
self, experiences: Tuple[torch.Tensor, ...], gamma: float
|
||||
) -> torch.Tensor:
|
||||
def _get_critic_loss(self, experiences, gamma):
|
||||
"""Return element-wise critic loss."""
|
||||
states, actions, rewards, next_states, dones = experiences[:5]
|
||||
|
||||
@@ -99,9 +98,7 @@ class Agent(TD3Agent):
|
||||
torch.cat((next_states, next_actions), dim=-1)
|
||||
)
|
||||
target_values = torch.min(target_values1, target_values2)
|
||||
target_values = (
|
||||
rewards + (self.hyper_params["GAMMA"] * target_values * masks).detach()
|
||||
)
|
||||
target_values = rewards + (gamma * target_values * masks).detach()
|
||||
|
||||
# train critic
|
||||
values1 = self.critic1(torch.cat((states, actions), dim=-1))
|
||||
@@ -113,19 +110,7 @@ class Agent(TD3Agent):
|
||||
return critic1_loss_element_wise, critic2_loss_element_wise
|
||||
|
||||
# pylint: disable=too-many-statements
|
||||
def update_model(
|
||||
self,
|
||||
experiences: Tuple[
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
torch.Tensor,
|
||||
List[int],
|
||||
],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
def update_model(self, experiences):
|
||||
"""Train the model after each episode."""
|
||||
states, actions, rewards, next_states, dones, weights, indices, eps_d = (
|
||||
experiences
|
||||
@@ -195,7 +180,7 @@ class Agent(TD3Agent):
|
||||
def pretrain(self):
|
||||
"""Pretraining steps."""
|
||||
pretrain_loss = list()
|
||||
print("[INFO] Pre-Train %d steps." % self.hyper_params["PRETRAIN_STEP"])
|
||||
print ("[INFO] Pre-Train %d steps." % self.hyper_params["PRETRAIN_STEP"])
|
||||
for i_step in range(1, self.hyper_params["PRETRAIN_STEP"] + 1):
|
||||
loss = self.update_model()
|
||||
pretrain_loss.append(loss) # for logging
|
||||
|
||||
Reference in New Issue
Block a user