pulsatrix
Loading...
Searching...
No Matches
dqn_loss.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8#include <vector>
9
11#include "pulsatrix/tensor.hpp"
12
13namespace pulsatrix {
14
37class DQNLoss {
38public:
43 explicit DQNLoss(DeviceBackend* backend);
44
64 [[nodiscard]] float forward(const Tensor& q_values, const Tensor& actions, const Tensor& targets);
65
78 [[nodiscard]] Tensor backward() const;
79
80private:
81 DeviceBackend* backend_;
82 Tensor last_q_values_;
83 Tensor last_targets_;
87 Tensor last_action_indices_;
88 bool has_forwarded_ = false;
89};
90
91} // namespace pulsatrix
loss = mean_b( (q_values[b, a_b] - targets[b,0])^2 ), where a_b is the action actually taken on trans...
Definition dqn_loss.hpp:37
Tensor backward() const
Gradient w.r.t. q_values: 2*(q_values[b,a_b] - targets[b,0])/N in the taken action's column,...
float forward(const Tensor &q_values, const Tensor &actions, const Tensor &targets)
Computes the masked MSE and caches everything backward() needs.
DQNLoss(DeviceBackend *backend)
Constructs a DQN loss.
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
Definition acquisition_functions.hpp:16
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).