pulsatrix
Loading...
Searching...
No Matches
hypergrid_env.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8#include <vector>
9
12#include "pulsatrix/tensor.hpp"
13
14namespace pulsatrix {
15
31class HyperGridEnv : public Environment {
32public:
49 explicit HyperGridEnv(DeviceBackend* backend, int64_t ndim = 2, int64_t side_length = 8, float r0 = 0.1f,
50 float r1 = 0.5f, float r2 = 2.0f, int64_t max_steps = 64);
51
57 [[nodiscard]] Tensor reset() override;
58
68 [[nodiscard]] Tensor reset(const Tensor& initial_state) override;
69
88 [[nodiscard]] StepResult step(const Tensor& action) override;
89
102 [[nodiscard]] float reward(const Tensor& state) const;
103
120 [[nodiscard]] float backward_log_prob(const Tensor& state, int64_t action) const;
121
134 [[nodiscard]] std::vector<bool> valid_actions_mask(const Tensor& state) const;
135
137 [[nodiscard]] int64_t observation_dim() const override { return ndim_; }
138
140 [[nodiscard]] int64_t action_dim() const override { return ndim_ + 1; }
141
143 [[nodiscard]] bool is_discrete() const override { return true; }
144
147 [[nodiscard]] int64_t ndim() const { return ndim_; }
148
150 [[nodiscard]] int64_t side_length() const { return side_length_; }
151
153 [[nodiscard]] int64_t max_steps() const { return max_steps_; }
154
156 [[nodiscard]] int64_t step_count() const { return step_count_; }
157
158private:
160 [[nodiscard]] Tensor observation() const;
161
164 void decode_state(const Tensor& state, std::vector<int64_t>& out) const;
165
166 DeviceBackend* backend_;
167 int64_t ndim_;
168 int64_t side_length_;
169 float r0_;
170 float r1_;
171 float r2_;
172 int64_t max_steps_;
173
174 std::vector<int64_t> coord_;
175 int64_t step_count_ = 0;
176 bool has_reset_ = false;
177 bool done_ = false;
178};
179
180} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Base class for every RL environment (CartPoleEnv, and whatever later phases add).
Definition environment.hpp:52
An n-dimensional grid: state is an integer coordinate in [0, H-1]^ndim, actions increment one coordin...
Definition hypergrid_env.hpp:31
HyperGridEnv(DeviceBackend *backend, int64_t ndim=2, int64_t side_length=8, float r0=0.1f, float r1=0.5f, float r2=2.0f, int64_t max_steps=64)
Constructs a fresh, not-yet-reset HyperGrid environment.
bool is_discrete() const override
True – HyperGrid's action space is discrete.
Definition hypergrid_env.hpp:143
StepResult step(const Tensor &action) override
Advances the episode one step: increments a coordinate, or stops.
float reward(const Tensor &state) const
The reward of an arbitrary valid grid state, independent of this environment's current episode state.
std::vector< bool > valid_actions_mask(const Tensor &state) const
Which actions are legal to take from an arbitrary valid state.
Tensor reset(const Tensor &initial_state) override
Starts a new episode from an exact caller-supplied state.
Tensor reset() override
Starts a new episode at the origin (0, ..., 0).
int64_t side_length() const
Number of cells per dimension (H), as passed to the constructor.
Definition hypergrid_env.hpp:150
int64_t max_steps() const
Episode length safety cap this environment was constructed with.
Definition hypergrid_env.hpp:153
float backward_log_prob(const Tensor &state, int64_t action) const
The closed-form uniform backward-policy log-probability P_B(a|state).
int64_t action_dim() const override
ndim() + 1 – one increment action per dimension, plus stop.
Definition hypergrid_env.hpp:140
int64_t ndim() const
Number of grid dimensions. Alias for observation_dim(), named for readability at HyperGrid-specific c...
Definition hypergrid_env.hpp:147
int64_t step_count() const
Steps taken since the last reset().
Definition hypergrid_env.hpp:156
int64_t observation_dim() const override
Number of grid dimensions, as passed to the constructor.
Definition hypergrid_env.hpp:137
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 RL environment interface (gymnasium-shaped reset/step) + StepResult.
Definition acquisition_functions.hpp:16
What one Environment::step() produces: the next observation, this step's reward, and whether the epis...
Definition environment.hpp:25
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).