mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Update logger
This commit is contained in:
@@ -1,10 +1,6 @@
|
||||
from .config import *
|
||||
from .normalizer import *
|
||||
from .misc import *
|
||||
from .tf_logger import Logger
|
||||
from .logger import get_logger
|
||||
from .plot import Plotter
|
||||
from .schedule import *
|
||||
import logging
|
||||
logging.basicConfig(format='%(asctime)s - %(name)s - %(levelname)s: %(message)s')
|
||||
logger = logging.getLogger('MAIN')
|
||||
logger.setLevel(logging.INFO)
|
||||
@@ -0,0 +1,73 @@
|
||||
#######################################################################
|
||||
# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) #
|
||||
# Permission given to modify the code as long as you keep this #
|
||||
# declaration at the top #
|
||||
#######################################################################
|
||||
|
||||
from tensorboardX import SummaryWriter
|
||||
import os
|
||||
import numpy as np
|
||||
import torch
|
||||
import logging
|
||||
logging.basicConfig(format='%(asctime)s - %(name)s - %(levelname)s: %(message)s')
|
||||
from .misc import *
|
||||
|
||||
def get_logger(name='MAIN', file_name=None, log_dir='./log', skip=False, level=logging.INFO):
|
||||
logger = logging.getLogger(name)
|
||||
logger.setLevel(level)
|
||||
if file_name is not None:
|
||||
file_name = '%s-%s' % (file_name, get_time_str())
|
||||
fh = logging.FileHandler('%s/%s.txt' % (log_dir, file_name))
|
||||
fh.setFormatter(logging.Formatter('%(asctime)s - %(name)s - %(levelname)s: %(message)s'))
|
||||
fh.setLevel(level)
|
||||
logger.addHandler(fh)
|
||||
return Logger(log_dir, logger, skip)
|
||||
|
||||
class Logger(object):
|
||||
def __init__(self, log_dir, vanilla_logger, skip=False):
|
||||
try:
|
||||
for f in os.listdir(log_dir):
|
||||
if not f.startswith('events'):
|
||||
continue
|
||||
os.remove('%s/%s' % (log_dir, f))
|
||||
except IOError:
|
||||
os.mkdir(log_dir)
|
||||
if not skip:
|
||||
self.writer = SummaryWriter(log_dir)
|
||||
self.info = vanilla_logger.info
|
||||
self.debug = vanilla_logger.debug
|
||||
self.warning = vanilla_logger.warning
|
||||
self.skip = skip
|
||||
self.all_steps = {}
|
||||
|
||||
def to_numpy(self, v):
|
||||
if isinstance(v, torch.autograd.Variable):
|
||||
v = v.data
|
||||
if isinstance(v, torch.FloatTensor):
|
||||
v = v.cpu().detach().numpy()
|
||||
return v
|
||||
|
||||
def get_step(self, tag):
|
||||
if tag not in self.all_steps:
|
||||
self.all_steps[tag] = 0
|
||||
step = self.all_steps[tag]
|
||||
self.all_steps[tag] += 1
|
||||
return step
|
||||
|
||||
def scalar_summary(self, tag, value, step=None):
|
||||
if self.skip:
|
||||
return
|
||||
self.to_numpy(value)
|
||||
if step is None:
|
||||
step = self.get_step(tag)
|
||||
if np.isscalar(value):
|
||||
value = np.asarray([value])
|
||||
self.writer.add_scalar(tag, value, step)
|
||||
|
||||
def histo_summary(self, tag, values, step=None):
|
||||
if self.skip:
|
||||
return
|
||||
self.to_numpy(values)
|
||||
if step is None:
|
||||
step = self.get_step(tag)
|
||||
self.writer.add_histogram(tag, values, step, bins=1000)
|
||||
@@ -120,3 +120,6 @@ class Batcher:
|
||||
indices = np.arange(self.num_entries)
|
||||
np.random.shuffle(indices)
|
||||
self.data = [d[indices] for d in self.data]
|
||||
|
||||
# def torch_max(tensor, dim):
|
||||
# return torch.max(tensor, dim=dim, keepdim=True)[0]
|
||||
|
||||
+19
-21
@@ -21,7 +21,7 @@ def dqn_cart_pole():
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 200
|
||||
config.exploration_steps = 1000
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.double_q = True
|
||||
# config.double_q = False
|
||||
run_episodes(DQNAgent(config))
|
||||
@@ -39,7 +39,7 @@ def a2c_cart_pole():
|
||||
config.network_fn = lambda state_dim, action_dim: ActorCriticNet(action_dim, FCBody(state_dim))
|
||||
config.policy_fn = SamplePolicy
|
||||
config.discount = 0.99
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.gae_tau = 1.0
|
||||
config.entropy_weight = 0.01
|
||||
config.rollout_length = 5
|
||||
@@ -58,7 +58,7 @@ def categorical_dqn_cart_pole():
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 200
|
||||
config.exploration_steps = 100
|
||||
config.logger = Logger('./log', logger, skip=True)
|
||||
config.logger = get_logger(skip=True)
|
||||
config.categorical_v_max = 100
|
||||
config.categorical_v_min = -100
|
||||
config.categorical_n_atoms = 50
|
||||
@@ -76,7 +76,7 @@ def quantile_regression_dqn_cart_pole():
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 200
|
||||
config.exploration_steps = 100
|
||||
config.logger = Logger('./log', logger, skip=True)
|
||||
config.logger = get_logger(skip=True)
|
||||
config.num_quantiles = 20
|
||||
run_episodes(QuantileRegressionDQNAgent(config))
|
||||
|
||||
@@ -92,7 +92,7 @@ def n_step_dqn_cart_pole():
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 200
|
||||
config.rollout_length = 5
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
run_iterations(NStepDQNAgent(config))
|
||||
|
||||
def ppo_cart_pole():
|
||||
@@ -105,7 +105,7 @@ def ppo_cart_pole():
|
||||
config.network_fn = lambda state_dim, action_dim: \
|
||||
CategoricalActorCriticWrapper(state_dim, action_dim, network_fn, optimizer_fn)
|
||||
config.discount = 0.99
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.use_gae = True
|
||||
config.gae_tau = 0.95
|
||||
config.entropy_weight = 0.01
|
||||
@@ -132,7 +132,7 @@ def option_critic_cart_pole():
|
||||
config.rollout_length = 5
|
||||
config.termination_regularizer = 0.01
|
||||
config.entropy_weight = 0.01
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
run_iterations(OptionCriticAgent(config))
|
||||
|
||||
## Atari games
|
||||
@@ -152,7 +152,7 @@ def dqn_pixel_atari(name):
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 10000
|
||||
config.exploration_steps= 50000
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
# config.double_q = True
|
||||
config.double_q = False
|
||||
run_episodes(DQNAgent(config))
|
||||
@@ -175,7 +175,7 @@ def a2c_pixel_atari(name):
|
||||
config.entropy_weight = 0.01
|
||||
config.rollout_length = 5
|
||||
config.gradient_clip = 0.5
|
||||
config.logger = Logger('./log', logger, skip=True)
|
||||
config.logger = get_logger(skip=True)
|
||||
run_iterations(A2CAgent(config))
|
||||
|
||||
def categorical_dqn_pixel_atari(name):
|
||||
@@ -193,7 +193,7 @@ def categorical_dqn_pixel_atari(name):
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
config.target_network_update_freq = 10000
|
||||
config.exploration_steps= 50000
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.double_q = False
|
||||
config.categorical_v_max = 10
|
||||
config.categorical_v_min = -10
|
||||
@@ -215,7 +215,7 @@ def quantile_regression_dqn_pixel_atari(name):
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 10000
|
||||
config.exploration_steps= 50000
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.double_q = False
|
||||
config.num_quantiles = 200
|
||||
run_episodes(QuantileRegressionDQNAgent(config))
|
||||
@@ -236,7 +236,7 @@ def n_step_dqn_pixel_atari(name):
|
||||
config.target_network_update_freq = 10000
|
||||
config.rollout_length = 5
|
||||
config.gradient_clip = 5
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
run_iterations(NStepDQNAgent(config))
|
||||
|
||||
def ppo_pixel_atari(name):
|
||||
@@ -253,7 +253,7 @@ def ppo_pixel_atari(name):
|
||||
config.state_normalizer = ImageNormalizer()
|
||||
config.reward_normalizer = SignNormalizer()
|
||||
config.discount = 0.99
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.use_gae = True
|
||||
config.gae_tau = 0.95
|
||||
config.entropy_weight = 0.01
|
||||
@@ -284,7 +284,7 @@ def option_ciritc_pixel_atari(name):
|
||||
config.max_steps = 1e8
|
||||
config.entropy_weight = 0.01
|
||||
config.termination_regularizer = 0.01
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
run_iterations(OptionCriticAgent(config))
|
||||
|
||||
def dqn_ram_atari(name):
|
||||
@@ -301,7 +301,7 @@ def dqn_ram_atari(name):
|
||||
config.target_network_update_freq = 10000
|
||||
config.max_episode_length = 0
|
||||
config.exploration_steps= 100
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
config.double_q = True
|
||||
# config.double_q = False
|
||||
run_episodes(DQNAgent(config))
|
||||
@@ -333,7 +333,7 @@ def ppo_continuous():
|
||||
config.num_mini_batches = 32
|
||||
config.ppo_ratio_clip = 0.2
|
||||
config.iteration_log_interval = 1
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
run_iterations(PPOAgent(config))
|
||||
|
||||
def ddpg_continuous():
|
||||
@@ -360,7 +360,7 @@ def ddpg_continuous():
|
||||
config.random_process_fn = lambda action_dim: GaussianProcess(action_dim, LinearSchedule(0.3, 0, 1e6))
|
||||
config.min_memory_size = 64
|
||||
config.target_network_mix = 1e-3
|
||||
config.logger = Logger('./log', logger)
|
||||
config.logger = get_logger()
|
||||
run_episodes(DDPGAgent(config))
|
||||
|
||||
def plot():
|
||||
@@ -397,11 +397,9 @@ if __name__ == '__main__':
|
||||
mkdir('dataset')
|
||||
mkdir('log')
|
||||
set_one_thread()
|
||||
# logger.setLevel(logging.DEBUG)
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
# dqn_cart_pole()
|
||||
# a2c_cart_pole()
|
||||
a2c_cart_pole()
|
||||
# categorical_dqn_cart_pole()
|
||||
# quantile_regression_dqn_cart_pole()
|
||||
# n_step_dqn_cart_pole()
|
||||
@@ -417,7 +415,7 @@ if __name__ == '__main__':
|
||||
# option_ciritc_pixel_atari('BreakoutNoFrameskip-v4')
|
||||
# dqn_ram_atari('Breakout-ramNoFrameskip-v4')
|
||||
|
||||
ddpg_continuous()
|
||||
# ddpg_continuous()
|
||||
# ppo_continuous()
|
||||
|
||||
# action_conditional_video_prediction()
|
||||
|
||||
Reference in New Issue
Block a user