diff --git a/.gitignore b/.gitignore index 5798e31..220aa96 100644 --- a/.gitignore +++ b/.gitignore @@ -8,6 +8,8 @@ exp_* upload.py *.sh data +draw_* +log # C extensions *.so diff --git a/common.py b/common.py deleted file mode 100644 index 24cec97..0000000 --- a/common.py +++ /dev/null @@ -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 \ No newline at end of file diff --git a/dqn_agent.py b/dqn_agent.py index 510fabb..c2436e8 100644 --- a/dqn_agent.py +++ b/dqn_agent.py @@ -10,9 +10,9 @@ from policy import * import numpy as np class DQNAgent: - def __init__(self, task, network_fn, policy_fn, replay_fn, discount, step_limit, target_network_update_freq): - self.learning_network = network_fn() - self.target_network = network_fn() + def __init__(self, task, network_fn, optimizer_fn, policy_fn, replay_fn, discount, step_limit, target_network_update_freq): + self.learning_network = network_fn(optimizer_fn) + self.target_network = network_fn(optimizer_fn) self.task = task self.step_limit = step_limit self.replay = replay_fn() diff --git a/network.py b/network.py index 5cf05d4..0480771 100644 --- a/network.py +++ b/network.py @@ -9,17 +9,15 @@ from torch.autograd import Variable import torch.nn as nn import torch.nn.functional as F import numpy as np -from SMDWrapper import SMDWrapper class FullyConnectedNet(nn.Module): - def __init__(self, dims, learning_rate, gpu=True): + def __init__(self, dims, optimizer_fn, gpu=True): super(FullyConnectedNet, self).__init__() self.fc1 = nn.Linear(dims[0], dims[1]) self.fc2 = nn.Linear(dims[1], dims[2]) self.fc3 = nn.Linear(dims[2], dims[3]) self.criterion = nn.MSELoss() - self.learning_rate = learning_rate - self.optimizer = torch.optim.SGD(self.parameters(), learning_rate) + self.optimizer = optimizer_fn(self.parameters()) self.gpu = gpu and torch.cuda.is_available() if self.gpu: print 'Transferring network to GPU...' @@ -58,13 +56,3 @@ class FullyConnectedNet(nn.Module): def output_transfer(self, 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() diff --git a/task.py b/task.py index 98dc6b8..fbc8d69 100644 --- a/task.py +++ b/task.py @@ -7,6 +7,7 @@ import gym import sys from dqn_agent import * +import torch.optim class BasicTask: def transfer_state(self, state): @@ -32,8 +33,8 @@ class MountainCar(BasicTask): def __init__(self): self.env = gym.make(self.name) self.env._max_episode_steps = sys.maxsize - - self.network_fn = lambda learning_rate=0.01: FullyConnectedNet([self.state_space_size, 50, 200, self.action_space_size], learning_rate) + self.optimizer_fn = lambda params: torch.optim.SGD(params, 0.001) + 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.replay_fn = lambda: Replay(memory_size=10000, batch_size=10) @@ -48,19 +49,15 @@ class CartPole(BasicTask): def __init__(self): self.env = gym.make(self.name) - - self.network_fn = lambda learning_rate=0.01: FullyConnectedNet([self.state_space_size, 50, 200, self.action_space_size], learning_rate) - self.policy_fn = lambda: GreedyPolicy(epsilon=0.5, decay_factor=0.95, min_epsilon=0.1) + self.optimizer_fn = lambda params: torch.optim.SGD(params, 0.001) + 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.99, min_epsilon=0.01) self.replay_fn = lambda: Replay(memory_size=10000, batch_size=10) if __name__ == '__main__': - task = MountainCar() - bp_network_fn = lambda learning_rate=0.001: FullyConnectedNet([task.state_space_size, 50, 200, task.action_space_size], learning_rate, gpu=False) - def smd_network_fn(learning_rate=0.001): - bp_network = bp_network_fn(learning_rate) - return SMDNetworkWrapper(bp_network) - - agent = DQNAgent(task, smd_network_fn, task.policy_fn, task.replay_fn, + task = CartPole() + optimizer_fn = lambda params: torch.optim.SGD(params, 0.001) + agent = DQNAgent(task, task.network_fn, optimizer_fn, task.policy_fn, task.replay_fn, task.discount, task.step_limit, task.target_network_update_freq) window_size = 100 ep = 0