mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-11 11:53:01 +08:00
Major update
This commit is contained in:
@@ -7,6 +7,7 @@ import numpy as np
|
||||
import torch
|
||||
from torch.autograd import Variable
|
||||
import torch.nn as nn
|
||||
from utils import *
|
||||
|
||||
class OneStepQLearning:
|
||||
def __init__(self, config, learning_network, target_network):
|
||||
@@ -60,11 +61,7 @@ class OneStepQLearning:
|
||||
self.optimizer.zero_grad()
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
|
||||
for param, worker_param in zip(
|
||||
self.learning_network.parameters(), self.worker_network.parameters()):
|
||||
if param.grad is not None:
|
||||
break
|
||||
param._grad = worker_param.grad
|
||||
sync_grad(self.learning_network, self.worker_network)
|
||||
self.optimizer.step()
|
||||
self.worker_network.load_state_dict(self.learning_network.state_dict())
|
||||
self.worker_network.reset(terminal)
|
||||
|
||||
Reference in New Issue
Block a user