mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-07 16:40:41 +08:00
Minor update for DQN
This commit is contained in:
+9
-13
@@ -16,29 +16,25 @@ import torch
|
||||
class DQNAgent:
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.learning_network = config.network_fn(config.optimizer_fn)
|
||||
self.target_network = config.network_fn(config.optimizer_fn)
|
||||
self.learning_network = config.network_fn()
|
||||
self.target_network = config.network_fn()
|
||||
self.optimizer = config.optimizer_fn(self.learning_network.parameters())
|
||||
self.criterion = nn.MSELoss()
|
||||
self.target_network.load_state_dict(self.learning_network.state_dict())
|
||||
self.task = config.task_fn()
|
||||
self.replay = config.replay_fn()
|
||||
self.policy = config.policy_fn()
|
||||
self.total_steps = 0
|
||||
self.history_buffer = None
|
||||
|
||||
def episode(self, deterministic=False):
|
||||
episode_start_time = time.time()
|
||||
state = self.task.reset()
|
||||
if self.history_buffer is None:
|
||||
self.history_buffer = [np.zeros_like(state)] * self.config.history_length
|
||||
else:
|
||||
self.history_buffer.pop(0)
|
||||
self.history_buffer.append(state)
|
||||
self.history_buffer = [state] * self.config.history_length
|
||||
state = np.vstack(self.history_buffer)
|
||||
total_reward = 0.0
|
||||
steps = 0
|
||||
while True:
|
||||
value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), False)
|
||||
value = value.cpu().data.numpy().flatten()
|
||||
value = self.learning_network.predict(np.stack([self.task.normalize_state(state)]), True).flatten()
|
||||
if deterministic:
|
||||
action = np.argmax(value)
|
||||
elif self.total_steps < self.config.exploration_steps:
|
||||
@@ -97,10 +93,10 @@ class DQNAgent:
|
||||
actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1)
|
||||
q = self.learning_network.predict(states, False)
|
||||
q = q.gather(1, actions).squeeze(1)
|
||||
loss = self.learning_network.criterion(q, q_next)
|
||||
self.learning_network.zero_grad()
|
||||
loss = self.criterion(q, q_next)
|
||||
self.optimizer.zero_grad()
|
||||
loss.backward()
|
||||
self.learning_network.optimizer.step()
|
||||
self.optimizer.step()
|
||||
if not deterministic and self.total_steps % self.config.target_network_update_freq == 0:
|
||||
self.target_network.load_state_dict(self.learning_network.state_dict())
|
||||
if not deterministic and self.total_steps > self.config.exploration_steps:
|
||||
|
||||
@@ -14,7 +14,7 @@ def dqn_cart_pole():
|
||||
config = Config()
|
||||
config.task_fn = lambda: CartPole()
|
||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, 0.001)
|
||||
config.network_fn = lambda optimizer_fn: FCNet([8, 50, 200, 2], optimizer_fn)
|
||||
config.network_fn = lambda: FCNet([8, 50, 200, 2])
|
||||
# config.network_fn = lambda optimizer_fn: DuelingFCNet([8, 50, 200, 2], optimizer_fn)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: Replay(memory_size=10000, batch_size=10)
|
||||
@@ -75,7 +75,7 @@ def dqn_pixel_atari(name):
|
||||
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False)
|
||||
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 optimizer_fn: NatureConvNet(config.history_length, action_dim, optimizer_fn)
|
||||
config.network_fn = lambda: NatureConvNet(config.history_length, action_dim)
|
||||
# config.network_fn = lambda optimizer_fn: DuelingNatureConvNet(config.history_length, n_actions, optimizer_fn)
|
||||
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)
|
||||
@@ -141,8 +141,7 @@ def dqn_fruit():
|
||||
config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9)
|
||||
config.reward_weight = np.ones(10) / 10
|
||||
config.hybrid_reward = False
|
||||
config.network_fn = lambda optimizer_fn: FruitHRFCNet(
|
||||
98, 4, config.reward_weight, optimizer_fn)
|
||||
config.network_fn = lambda: FruitHRFCNet(98, 4, config.reward_weight)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: Replay(memory_size=10000, batch_size=15)
|
||||
config.discount = 0.95
|
||||
@@ -163,8 +162,7 @@ def hrdqn_fruit():
|
||||
config.hybrid_reward = True
|
||||
config.reward_weight = np.ones(10) / 10
|
||||
config.optimizer_fn = lambda params: torch.optim.SGD(params, 0.01, momentum=0.9)
|
||||
config.network_fn = lambda optimizer_fn: FruitHRFCNet(
|
||||
98, 4, config.reward_weight, optimizer_fn)
|
||||
config.network_fn = lambda optimizer_fn: FruitHRFCNet(98, 4, config.reward_weight)
|
||||
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1)
|
||||
config.replay_fn = lambda: HybridRewardReplay(memory_size=10000, batch_size=15)
|
||||
config.discount = 0.95
|
||||
@@ -280,7 +278,7 @@ if __name__ == '__main__':
|
||||
# logger.setLevel(logging.DEBUG)
|
||||
logger.setLevel(logging.INFO)
|
||||
|
||||
dqn_cart_pole()
|
||||
# dqn_cart_pole()
|
||||
# async_cart_pole()
|
||||
# a3c_cart_pole()
|
||||
# a3c_continuous()
|
||||
@@ -290,7 +288,7 @@ if __name__ == '__main__':
|
||||
# dqn_fruit()
|
||||
# hrdqn_fruit()
|
||||
|
||||
# dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
dqn_pixel_atari('PongNoFrameskip-v4')
|
||||
# async_pixel_atari('PongNoFrameskip-v4')
|
||||
# a3c_pixel_atari('PongNoFrameskip-v4')
|
||||
|
||||
|
||||
@@ -13,8 +13,6 @@ import numpy as np
|
||||
# Base class for all kinds of network
|
||||
class BasicNet:
|
||||
def __init__(self, optimizer_fn, gpu, LSTM=False):
|
||||
if optimizer_fn is not None:
|
||||
self.optimizer = optimizer_fn(self.parameters())
|
||||
self.gpu = gpu and torch.cuda.is_available()
|
||||
self.LSTM = LSTM
|
||||
if self.gpu:
|
||||
|
||||
@@ -15,8 +15,7 @@ class NatureConvNet(nn.Module, VanillaNet):
|
||||
self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1)
|
||||
self.fc4 = nn.Linear(7 * 7 * 64, 512)
|
||||
self.fc5 = nn.Linear(512, n_actions)
|
||||
self.criterion = nn.MSELoss()
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
BasicNet.__init__(self, None, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
@@ -37,8 +36,7 @@ class DuelingNatureConvNet(nn.Module, DuelingNet):
|
||||
self.fc4 = nn.Linear(7 * 7 * 64, 512)
|
||||
self.fc_advantage = nn.Linear(512, n_actions)
|
||||
self.fc_value = nn.Linear(512, 1)
|
||||
self.criterion = nn.MSELoss()
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
BasicNet.__init__(self, None, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
|
||||
@@ -13,7 +13,6 @@ class FCNet(nn.Module, VanillaNet):
|
||||
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()
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
@@ -32,7 +31,6 @@ class DuelingFCNet(nn.Module, DuelingNet):
|
||||
self.fc2 = nn.Linear(dims[1], dims[2])
|
||||
self.fc_value = nn.Linear(dims[2], 1)
|
||||
self.fc_advantage = nn.Linear(dims[2], dims[3])
|
||||
self.criterion = nn.MSELoss()
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
@@ -67,7 +65,6 @@ class FruitHRFCNet(nn.Module, VanillaNet):
|
||||
hidden_size = 250
|
||||
self.fc1 = nn.Linear(state_dim, hidden_size)
|
||||
self.fc2 = nn.ModuleList([nn.Linear(hidden_size, action_dim) for _ in head_weights])
|
||||
self.criterion = nn.MSELoss()
|
||||
self.head_weights = head_weights
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
|
||||
@@ -93,7 +90,6 @@ class FruitMultiStatesFCNet(nn.Module, BasicNet):
|
||||
hidden_size = 250
|
||||
self.fc1 = nn.ModuleList([nn.Linear(state_dim, hidden_size) for _ in head_weights])
|
||||
self.fc2 = nn.ModuleList([nn.Linear(hidden_size, action_dim) for _ in head_weights])
|
||||
self.criterion = nn.MSELoss()
|
||||
self.head_weights = head_weights
|
||||
self.state_dim = state_dim
|
||||
self.n_heads = head_weights.shape[0]
|
||||
|
||||
Reference in New Issue
Block a user