|
pulsatrix
|
DQN Bellman target computation (vanilla + Double DQN) and target-network hard sync. More...
#include "pulsatrix/device_backend.hpp"#include "pulsatrix/module.hpp"#include "pulsatrix/tensor.hpp"
Go to the source code of this file.
Namespaces | |
| namespace | pulsatrix |
Functions | |
| Tensor | pulsatrix::ComputeDQNTarget (const Tensor &next_q_target, const Tensor &rewards, const Tensor &dones, float gamma, DeviceBackend *backend) |
Vanilla DQN Bellman target (Mnih et al. 2015): targets[b,0] = rewards[b,0] + gamma * (1 - dones[b,0]) * max_a next_q_target[b,a]. | |
| Tensor | pulsatrix::ComputeDoubleDQNTarget (const Tensor &next_q_online, const Tensor &next_q_target, const Tensor &rewards, const Tensor &dones, float gamma, DeviceBackend *backend) |
Double DQN Bellman target (van Hasselt et al. 2016, arXiv:1509.06461): a* = argmax_a next_q_online[b,a], then targets[b,0] = rewards[b,0] + gamma * (1 - dones[b,0]) * next_q_target[b, a*]. | |
| void | pulsatrix::SyncTargetNetwork (Module &source, Module &destination) |
Hard target-network update: copies every parameter value of source into destination, element-wise and in place (Mnih et al. 2015's periodic full copy). | |
DQN Bellman target computation (vanilla + Double DQN) and target-network hard sync.