pulsatrix
Loading...
Searching...
No Matches
dqn_agent.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8
9#include "pulsatrix/agent.hpp"
11#include "pulsatrix/module.hpp"
12#include "pulsatrix/tensor.hpp"
13
14namespace pulsatrix {
15
37class DQNAgent : public Agent {
38public:
51 DQNAgent(Module* q_network, int64_t action_dim, float epsilon, DeviceBackend* backend, uint32_t seed = 42);
52
71 [[nodiscard]] Tensor act(const Tensor& observation) override;
72
87 [[nodiscard]] Tensor act_greedy(const Tensor& observation);
88
98 void set_epsilon(float epsilon);
99
101 [[nodiscard]] float epsilon() const { return epsilon_; }
102
104 [[nodiscard]] int64_t action_dim() const { return action_dim_; }
105
106private:
110 [[nodiscard]] float next_unit();
111
113 [[nodiscard]] int64_t next_index(int64_t bound);
114
116 [[nodiscard]] Tensor greedy_action(const Tensor& observation);
117
118 Module* q_network_;
119 int64_t action_dim_;
120 float epsilon_;
121 DeviceBackend* backend_;
122 uint32_t lcg_state_;
123};
124
125} // namespace pulsatrix
Abstract RL agent interface – the inference-time policy contract, act() only.
Base class for anything that maps an observation to an action.
Definition agent.hpp:24
The epsilon-greedy behaviour policy of Mnih et al. 2015: with probability epsilon act uniformly at ra...
Definition dqn_agent.hpp:37
Tensor act(const Tensor &observation) override
Chooses an action epsilon-greedily, advancing the internal LCG.
DQNAgent(Module *q_network, int64_t action_dim, float epsilon, DeviceBackend *backend, uint32_t seed=42)
Constructs an epsilon-greedy policy over a Q-network.
void set_epsilon(float epsilon)
Sets the exploration probability – the hook an epsilon-decay schedule drives.
Tensor act_greedy(const Tensor &observation)
The purely greedy (evaluation-time) policy: argmax with no exploration at all.
int64_t action_dim() const
Number of discrete actions, as passed to the constructor.
Definition dqn_agent.hpp:104
float epsilon() const
The current exploration probability.
Definition dqn_agent.hpp:101
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
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.
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).