mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-11 11:53:01 +08:00
Refactor interfaces
This commit is contained in:
@@ -8,6 +8,8 @@ exp_*
|
|||||||
upload.py
|
upload.py
|
||||||
*.sh
|
*.sh
|
||||||
data
|
data
|
||||||
|
draw_*
|
||||||
|
log
|
||||||
|
|
||||||
# C extensions
|
# C extensions
|
||||||
*.so
|
*.so
|
||||||
|
|||||||
@@ -1,56 +0,0 @@
|
|||||||
#######################################################################
|
|
||||||
# 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 #
|
|
||||||
#######################################################################
|
|
||||||
|
|
||||||
import tensorflow as tf
|
|
||||||
|
|
||||||
def fully_connected(model_name, layer_name, var_in, dim_in, dim_out,
|
|
||||||
initializer, transfer):
|
|
||||||
with tf.variable_scope(model_name):
|
|
||||||
with tf.variable_scope(layer_name):
|
|
||||||
W = tf.get_variable("W", [dim_in, dim_out],
|
|
||||||
initializer=initializer)
|
|
||||||
b = tf.get_variable("b", [dim_out],
|
|
||||||
initializer=initializer)
|
|
||||||
net = tf.nn.bias_add(tf.matmul(var_in, W), b)
|
|
||||||
phi = transfer(net)
|
|
||||||
return W, b, net, phi
|
|
||||||
|
|
||||||
class Relu:
|
|
||||||
def __init__(self):
|
|
||||||
self.gate_fun = tf.nn.relu
|
|
||||||
self.gate_fun_gradient = \
|
|
||||||
lambda phi, net: tf.where(net >= 0, tf.ones(tf.shape(net)), tf.zeros(tf.shape(net)))
|
|
||||||
|
|
||||||
|
|
||||||
class Tanh:
|
|
||||||
def __init__(self):
|
|
||||||
self.gate_fun = tf.tanh
|
|
||||||
self.gate_fun_gradient = \
|
|
||||||
lambda phi, net: tf.subtract(1.0, tf.pow(phi, 2))
|
|
||||||
|
|
||||||
class Identity:
|
|
||||||
def __init__(self):
|
|
||||||
self.gate_fun = tf.identity
|
|
||||||
self.gate_fun_gradient = \
|
|
||||||
lambda phi, net: tf.ones(tf.shape(phi))
|
|
||||||
|
|
||||||
def crossprop_layer(model_name, layer_name, var_in, dim_in, dim_hidden, dim_out, gate_fun, initializer):
|
|
||||||
with tf.variable_scope(model_name):
|
|
||||||
with tf.variable_scope(layer_name):
|
|
||||||
U = tf.get_variable('U', [dim_in, dim_hidden],
|
|
||||||
initializer=initializer)
|
|
||||||
b_hidden = tf.get_variable('b_hidden', [dim_hidden],
|
|
||||||
initializer=initializer)
|
|
||||||
W = tf.get_variable('W', [dim_hidden, dim_out],
|
|
||||||
initializer=initializer)
|
|
||||||
b_out = tf.get_variable('b_out', [dim_out],
|
|
||||||
initializer=initializer)
|
|
||||||
net = tf.matmul(var_in, U)
|
|
||||||
net = tf.nn.bias_add(net, b_hidden)
|
|
||||||
phi = gate_fun(net)
|
|
||||||
y = tf.matmul(phi, W)
|
|
||||||
y = tf.nn.bias_add(y, b_out)
|
|
||||||
return U, b_hidden, net, phi, W, b_out, y
|
|
||||||
+3
-3
@@ -10,9 +10,9 @@ from policy import *
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
class DQNAgent:
|
class DQNAgent:
|
||||||
def __init__(self, task, network_fn, policy_fn, replay_fn, discount, step_limit, target_network_update_freq):
|
def __init__(self, task, network_fn, optimizer_fn, policy_fn, replay_fn, discount, step_limit, target_network_update_freq):
|
||||||
self.learning_network = network_fn()
|
self.learning_network = network_fn(optimizer_fn)
|
||||||
self.target_network = network_fn()
|
self.target_network = network_fn(optimizer_fn)
|
||||||
self.task = task
|
self.task = task
|
||||||
self.step_limit = step_limit
|
self.step_limit = step_limit
|
||||||
self.replay = replay_fn()
|
self.replay = replay_fn()
|
||||||
|
|||||||
+2
-14
@@ -9,17 +9,15 @@ from torch.autograd import Variable
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from SMDWrapper import SMDWrapper
|
|
||||||
|
|
||||||
class FullyConnectedNet(nn.Module):
|
class FullyConnectedNet(nn.Module):
|
||||||
def __init__(self, dims, learning_rate, gpu=True):
|
def __init__(self, dims, optimizer_fn, gpu=True):
|
||||||
super(FullyConnectedNet, self).__init__()
|
super(FullyConnectedNet, self).__init__()
|
||||||
self.fc1 = nn.Linear(dims[0], dims[1])
|
self.fc1 = nn.Linear(dims[0], dims[1])
|
||||||
self.fc2 = nn.Linear(dims[1], dims[2])
|
self.fc2 = nn.Linear(dims[1], dims[2])
|
||||||
self.fc3 = nn.Linear(dims[2], dims[3])
|
self.fc3 = nn.Linear(dims[2], dims[3])
|
||||||
self.criterion = nn.MSELoss()
|
self.criterion = nn.MSELoss()
|
||||||
self.learning_rate = learning_rate
|
self.optimizer = optimizer_fn(self.parameters())
|
||||||
self.optimizer = torch.optim.SGD(self.parameters(), learning_rate)
|
|
||||||
self.gpu = gpu and torch.cuda.is_available()
|
self.gpu = gpu and torch.cuda.is_available()
|
||||||
if self.gpu:
|
if self.gpu:
|
||||||
print 'Transferring network to GPU...'
|
print 'Transferring network to GPU...'
|
||||||
@@ -58,13 +56,3 @@ class FullyConnectedNet(nn.Module):
|
|||||||
|
|
||||||
def output_transfer(self, y):
|
def output_transfer(self, y):
|
||||||
return y
|
return y
|
||||||
|
|
||||||
class SMDNetworkWrapper(SMDWrapper):
|
|
||||||
def __init__(self, net):
|
|
||||||
SMDWrapper.__init__(self, net)
|
|
||||||
|
|
||||||
def sync_with(self, src_net):
|
|
||||||
self.net.sync_with(src_net.net)
|
|
||||||
|
|
||||||
def parameters(self):
|
|
||||||
return self.net.parameters()
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@
|
|||||||
import gym
|
import gym
|
||||||
import sys
|
import sys
|
||||||
from dqn_agent import *
|
from dqn_agent import *
|
||||||
|
import torch.optim
|
||||||
|
|
||||||
class BasicTask:
|
class BasicTask:
|
||||||
def transfer_state(self, state):
|
def transfer_state(self, state):
|
||||||
@@ -32,8 +33,8 @@ class MountainCar(BasicTask):
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.env = gym.make(self.name)
|
self.env = gym.make(self.name)
|
||||||
self.env._max_episode_steps = sys.maxsize
|
self.env._max_episode_steps = sys.maxsize
|
||||||
|
self.optimizer_fn = lambda params: torch.optim.SGD(params, 0.001)
|
||||||
self.network_fn = lambda learning_rate=0.01: FullyConnectedNet([self.state_space_size, 50, 200, self.action_space_size], learning_rate)
|
self.network_fn = lambda optimizer_fn: FullyConnectedNet([self.state_space_size, 50, 200, self.action_space_size], optimizer_fn)
|
||||||
self.policy_fn = lambda: GreedyPolicy(epsilon=0.5, decay_factor=0.95, min_epsilon=0.1)
|
self.policy_fn = lambda: GreedyPolicy(epsilon=0.5, decay_factor=0.95, min_epsilon=0.1)
|
||||||
self.replay_fn = lambda: Replay(memory_size=10000, batch_size=10)
|
self.replay_fn = lambda: Replay(memory_size=10000, batch_size=10)
|
||||||
|
|
||||||
@@ -48,19 +49,15 @@ class CartPole(BasicTask):
|
|||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.env = gym.make(self.name)
|
self.env = gym.make(self.name)
|
||||||
|
self.optimizer_fn = lambda params: torch.optim.SGD(params, 0.001)
|
||||||
self.network_fn = lambda learning_rate=0.01: FullyConnectedNet([self.state_space_size, 50, 200, self.action_space_size], learning_rate)
|
self.network_fn = lambda optimizer_fn: FullyConnectedNet([self.state_space_size, 50, 200, self.action_space_size], optimizer_fn)
|
||||||
self.policy_fn = lambda: GreedyPolicy(epsilon=0.5, decay_factor=0.95, min_epsilon=0.1)
|
self.policy_fn = lambda: GreedyPolicy(epsilon=0.5, decay_factor=0.99, min_epsilon=0.01)
|
||||||
self.replay_fn = lambda: Replay(memory_size=10000, batch_size=10)
|
self.replay_fn = lambda: Replay(memory_size=10000, batch_size=10)
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
task = MountainCar()
|
task = CartPole()
|
||||||
bp_network_fn = lambda learning_rate=0.001: FullyConnectedNet([task.state_space_size, 50, 200, task.action_space_size], learning_rate, gpu=False)
|
optimizer_fn = lambda params: torch.optim.SGD(params, 0.001)
|
||||||
def smd_network_fn(learning_rate=0.001):
|
agent = DQNAgent(task, task.network_fn, optimizer_fn, task.policy_fn, task.replay_fn,
|
||||||
bp_network = bp_network_fn(learning_rate)
|
|
||||||
return SMDNetworkWrapper(bp_network)
|
|
||||||
|
|
||||||
agent = DQNAgent(task, smd_network_fn, task.policy_fn, task.replay_fn,
|
|
||||||
task.discount, task.step_limit, task.target_network_update_freq)
|
task.discount, task.step_limit, task.target_network_update_freq)
|
||||||
window_size = 100
|
window_size = 100
|
||||||
ep = 0
|
ep = 0
|
||||||
|
|||||||
Reference in New Issue
Block a user