mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Merge branch 'master' of https://github.com/ShangtongZhang/DeepRL
This commit is contained in:
+3
-2
@@ -17,7 +17,8 @@ PREFIX = '/local/data'
|
||||
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)
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False,
|
||||
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)
|
||||
@@ -26,7 +27,7 @@ def dqn_pixel_atari(name):
|
||||
config.discount = 0.99
|
||||
config.target_network_update_freq = 10000
|
||||
config.max_episode_length = 0
|
||||
config.exploration_steps= 50000
|
||||
config.exploration_steps = 50000
|
||||
config.logger = Logger('./log', logger)
|
||||
config.test_interval = 10
|
||||
config.test_repetitions = 1
|
||||
|
||||
@@ -99,7 +99,7 @@ def dqn_pixel_atari(name):
|
||||
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, n_actions)
|
||||
# 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.reward_shift_fn = lambda r: np.sign(r)
|
||||
|
||||
+5
-2
@@ -10,8 +10,11 @@ import numpy as np
|
||||
|
||||
class Logger(object):
|
||||
def __init__(self, log_dir, vanilla_logger, skip=False):
|
||||
for f in os.listdir(log_dir):
|
||||
os.remove('%s/%s' % (log_dir, f))
|
||||
try:
|
||||
for f in os.listdir(log_dir):
|
||||
os.remove('%s/%s' % (log_dir, f))
|
||||
except FileNotFoundError:
|
||||
os.mkdir(log_dir)
|
||||
if not skip:
|
||||
self.writer = SummaryWriter(log_dir)
|
||||
self.info = vanilla_logger.info
|
||||
|
||||
Reference in New Issue
Block a user