pulsatrix
Loading...
Searching...
No Matches
categorical_policy_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
41public:
53 CategoricalPolicyAgent(Module* policy_network, int64_t action_dim, DeviceBackend* backend, uint32_t seed = 42);
54
78 [[nodiscard]] Tensor act(const Tensor& observation) override;
79
97 [[nodiscard]] Tensor act_greedy(const Tensor& observation);
98
107 [[nodiscard]] float log_prob() const;
108
110 [[nodiscard]] int64_t action_dim() const { return action_dim_; }
111
112private:
116 [[nodiscard]] float next_unit();
117
119 [[nodiscard]] Tensor policy_logits(const Tensor& observation);
120
121 Module* policy_network_;
122 int64_t action_dim_;
123 DeviceBackend* backend_;
124 uint32_t lcg_state_;
125 float last_log_prob_ = 0.0f;
126 bool has_acted_ = false;
127};
128
129} // 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
Samples an action from softmax(policy_network(observation)) – the discrete-action stochastic policy e...
Definition categorical_policy_agent.hpp:40
Tensor act_greedy(const Tensor &observation)
The deterministic (evaluation-time) policy: argmax_a logits[a], no sampling.
Tensor act(const Tensor &observation) override
Samples an action from the policy's categorical distribution, advancing the LCG.
int64_t action_dim() const
Number of discrete actions, as passed to the constructor.
Definition categorical_policy_agent.hpp:110
float log_prob() const
The log-probability the policy assigned to the most recent act() sample.
CategoricalPolicyAgent(Module *policy_network, int64_t action_dim, DeviceBackend *backend, uint32_t seed=42)
Constructs a categorical policy over a logit-producing network.
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).