pulsatrix
Loading...
Searching...
No Matches
pulsatrix::HyperGridEnv Class Reference

An n-dimensional grid: state is an integer coordinate in [0, H-1]^ndim, actions increment one coordinate or stop the episode, reward is concentrated near the grid's corners. More...

#include <hypergrid_env.hpp>

Inheritance diagram for pulsatrix::HyperGridEnv:
Collaboration diagram for pulsatrix::HyperGridEnv:

Public Member Functions

 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.
 
Tensor reset () override
 Starts a new episode at the origin (0, ..., 0).
 
Tensor reset (const Tensor &initial_state) override
 Starts a new episode from an exact caller-supplied state.
 
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.
 
float backward_log_prob (const Tensor &state, int64_t action) const
 The closed-form uniform backward-policy log-probability P_B(a|state).
 
std::vector< bool > valid_actions_mask (const Tensor &state) const
 Which actions are legal to take from an arbitrary valid state.
 
int64_t observation_dim () const override
 Number of grid dimensions, as passed to the constructor.
 
int64_t action_dim () const override
 ndim() + 1 – one increment action per dimension, plus stop.
 
bool is_discrete () const override
 True – HyperGrid's action space is discrete.
 
int64_t ndim () const
 Number of grid dimensions. Alias for observation_dim(), named for readability at HyperGrid-specific call sites.
 
int64_t side_length () const
 Number of cells per dimension (H), as passed to the constructor.
 
int64_t max_steps () const
 Episode length safety cap this environment was constructed with.
 
int64_t step_count () const
 Steps taken since the last reset().
 
- Public Member Functions inherited from pulsatrix::Environment
virtual ~Environment ()=default
 

Detailed Description

An n-dimensional grid: state is an integer coordinate in [0, H-1]^ndim, actions increment one coordinate or stop the episode, reward is concentrated near the grid's corners.

Note
Deliberately not a reproduction of any single paper's exact reward constants – see mission_hypergrid_env.md's Design section for the reward shape this class implements and why it was chosen (a real, non-degenerate near-corner band distinct from the exact-corner set, hand-verifiable at named cells for the default side_length).
Deterministic reset() to the origin – unlike CartPoleEnv's randomized reset, this environment's whole purpose is a fixed, hand-traceable start state.
Reward is 0 at every non-terminal step and R(x) only at the terminal state (the stop action or the max_steps cap) – the standard GFlowNet convention that R is a property of the terminal object x, not of any intermediate transition.

Constructor & Destructor Documentation

◆ HyperGridEnv()

pulsatrix::HyperGridEnv::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 
)
explicit

Constructs a fresh, not-yet-reset HyperGrid environment.

Parameters
backendBackend to allocate observation tensors through. Not owned; must outlive this object.
ndimNumber of grid dimensions. Must be >= 1.
side_lengthNumber of cells per dimension (H). Must be >= 4 – below that the near-corner band (Design section) degenerates to exactly the corner set, which would make R1 and R2 indistinguishable rather than a real distinct band.
r0Base reward, everywhere. Must be >= 0.
r1Additional reward in the near-corner band. Must be >= 0.
r2Additional reward exactly at a corner. Must be >= 0.
max_stepsEpisode length safety cap – forces termination (as if stop were taken) if reached without one. Must be >= 1.
Exceptions
std::invalid_argumenton any violated precondition above – external boundary, same classification as CartPoleEnv's constructor.

Member Function Documentation

◆ action_dim()

int64_t pulsatrix::HyperGridEnv::action_dim ( ) const
inlineoverridevirtual

ndim() + 1 – one increment action per dimension, plus stop.

Implements pulsatrix::Environment.

◆ backward_log_prob()

float pulsatrix::HyperGridEnv::backward_log_prob ( const Tensor &  state,
int64_t  action 
) const

The closed-form uniform backward-policy log-probability P_B(a|state).

Parameters
stateGrid coordinate, shape (1, ndim()), each entry an integer-valued float in [0, side_length()-1].
actionDimension index [0, ndim()) – decrementing this coordinate must reach a valid parent state.
Returns
log(1 / count_nonzero(state)) – HyperGrid's action structure (exactly one way to reach any non-origin state, by incrementing one coordinate) makes the correct backward distribution uniform over state's nonzero coordinates; see mission_shared_gflownet_machinery.md's Recon for why this is closed-form rather than a second learned policy.
Exceptions
std::invalid_argumentif state is invalid (same validation as reward()), if action is out of [0, ndim()), or if state[action] == 0 (decrementing it would leave the grid – not a valid parent transition) – all external boundary.
Note
A pure function, like reward() – no precondition on reset() having been called.

◆ is_discrete()

bool pulsatrix::HyperGridEnv::is_discrete ( ) const
inlineoverridevirtual

True – HyperGrid's action space is discrete.

Implements pulsatrix::Environment.

◆ max_steps()

int64_t pulsatrix::HyperGridEnv::max_steps ( ) const
inline

Episode length safety cap this environment was constructed with.

◆ ndim()

int64_t pulsatrix::HyperGridEnv::ndim ( ) const
inline

Number of grid dimensions. Alias for observation_dim(), named for readability at HyperGrid-specific call sites.

◆ observation_dim()

int64_t pulsatrix::HyperGridEnv::observation_dim ( ) const
inlineoverridevirtual

Number of grid dimensions, as passed to the constructor.

Implements pulsatrix::Environment.

◆ reset() [1/2]

Tensor pulsatrix::HyperGridEnv::reset ( )
overridevirtual

Starts a new episode at the origin (0, ..., 0).

Returns
The initial observation, shape (1, ndim()).
Note
No randomness – HyperGrid's whole point is a fixed, hand-traceable start state.

Implements pulsatrix::Environment.

◆ reset() [2/2]

Tensor pulsatrix::HyperGridEnv::reset ( const Tensor &  initial_state)
overridevirtual

Starts a new episode from an exact caller-supplied state.

Parameters
initial_stateGrid coordinate, shape (1, ndim()), each entry an integer-valued float in [0, side_length()-1].
Returns
The initial observation (a copy of initial_state's values), shape (1, ndim()).
Exceptions
std::invalid_argumentif initial_state's shape is wrong, an entry is not within a small tolerance of an integer, or an entry is out of [0, side_length()-1] – all external boundary.

Implements pulsatrix::Environment.

◆ reward()

float pulsatrix::HyperGridEnv::reward ( const Tensor &  state) const

The reward of an arbitrary valid grid state, independent of this environment's current episode state.

Parameters
stateGrid coordinate, shape (1, ndim()), each entry an integer-valued float in [0, side_length()-1].
Returns
R(state) per the Design section's formula.
Exceptions
std::invalid_argumentif state's shape or entries are invalid – external boundary, same validation as reset(const Tensor&).
Note
A pure function – no precondition on reset() having been called, no mutation. Later missions (TrajectoryBalanceLoss and friends) need R(x) at a sampled terminal state directly, without re-driving the environment through step().

◆ side_length()

int64_t pulsatrix::HyperGridEnv::side_length ( ) const
inline

Number of cells per dimension (H), as passed to the constructor.

◆ step()

StepResult pulsatrix::HyperGridEnv::step ( const Tensor &  action)
overridevirtual

Advances the episode one step: increments a coordinate, or stops.

Parameters
actionShape (1, 1), holding the action index as a float: [0, ndim()) increment that coordinate, ndim() stops the episode at the current state.
Returns
The resulting observation, reward (0 unless this step is terminal), and whether the episode ended (via stop, an implicit max_steps cap, or neither).
Exceptions
std::invalid_argumentif reset() has never been called, the action's shape is wrong, the decoded index is out of range, or an increment action would move a coordinate past side_length()-1 (illegal off-grid move) – all external boundary: any of these can originate from an untrained or malformed policy.
std::logic_errorif the most recent step already reported done=true – a GFlowNet trajectory has no meaning past its terminal state; unlike CartPoleEnv, this is enforced here rather than left to the caller.
Note
Host boundary (GPU-native-kernels Mission 7): the grid walk is integer host logic. action (and every state argument of reset()/reward()/backward_log_prob()/ valid_actions_mask()) may live on any device – one device->host copy per call – and observations are returned through this environment's own backend.

Implements pulsatrix::Environment.

◆ step_count()

int64_t pulsatrix::HyperGridEnv::step_count ( ) const
inline

Steps taken since the last reset().

◆ valid_actions_mask()

std::vector< bool > pulsatrix::HyperGridEnv::valid_actions_mask ( const Tensor &  state) const

Which actions are legal to take from an arbitrary valid state.

Parameters
stateGrid coordinate, shape (1, ndim()), each entry an integer-valued float in [0, side_length()-1].
Returns
A mask of size action_dim(): entry i < ndim() is true iff state[i] < side_length()-1 (incrementing that coordinate stays on the grid); entry ndim() (stop) is always true.
Exceptions
std::invalid_argumentif state is invalid – same validation as reward().
Note
A pure function, like reward()/backward_log_prob() – no precondition on reset() having been called. GFlowNetForwardPolicy::sample() needs this mask to avoid ever sampling an illegal increment that would make step() throw.

The documentation for this class was generated from the following file: