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

Samples from softmax(mask(policy_network(observation))) – a categorical policy restricted to a caller-supplied set of valid actions at the current state. More...

#include <gflownet_forward_policy.hpp>

Public Member Functions

 GFlowNetForwardPolicy (Module *policy_network, int64_t action_dim, DeviceBackend *backend, uint32_t seed=42)
 Constructs a masked categorical policy over a logit-producing network.
 
GFlowNetSampledAction sample (const Tensor &observation, const std::vector< bool > &valid_actions)
 Samples an action from the policy's distribution, restricted to valid_actions.
 
std::vector< float > masked_probs (const Tensor &observation, const std::vector< bool > &valid_actions)
 The masked categorical distribution's probabilities, without sampling.
 
int64_t action_dim () const
 Number of actions, as passed to the constructor.
 

Detailed Description

Samples from softmax(mask(policy_network(observation))) – a categorical policy restricted to a caller-supplied set of valid actions at the current state.

Note
Byte-for-byte CategoricalPolicyAgent's LCG (Numerical Recipes constants) and numerically-stable log-softmax + inverse-CDF sampling, deliberately duplicated rather than reusing that class directly: CategoricalPolicyAgent::act() has no notion of a per-call action mask (every existing consumer – CartPole, etc. – has no invalid actions), and HyperGridEnv::step() throws on an off-grid increment, so a forward policy that ever sampled one would crash training. See mission_shared_gflownet_machinery.md's Recon for why this is a new, justified class rather than a modification to a closed, tested class from a different campaign.
Masking is the standard additive-mask technique: an invalid action's logit is driven to std::numeric_limits<float>::lowest() before softmax, giving it ~0 probability without risking a NaN from actual -infinity arithmetic.
Returns {action, log_prob} directly from sample() rather than CategoricalPolicyAgent's stateful act() + separate log_prob() accessor – a cleaner API this new class is free to choose, since nothing else depends on matching that older class's shape.
Host boundary (GPU-native-kernels Mission 7): masking, softmax and the inverse-CDF scan against the LCG draw are host logic. The network may run on any device; sample() and masked_probs() copy its (1, action_dim) logits to the host once per call, and the sampled action is returned through this policy's own backend.

Constructor & Destructor Documentation

◆ GFlowNetForwardPolicy()

pulsatrix::GFlowNetForwardPolicy::GFlowNetForwardPolicy ( Module *  policy_network,
int64_t  action_dim,
DeviceBackend *  backend,
uint32_t  seed = 42 
)

Constructs a masked categorical policy over a logit-producing network.

Parameters
policy_networkNetwork mapping a (1, observation_dim) observation to (1, action_dim) raw logits. Not owned; must outlive this object.
action_dimNumber of actions. Must be >= 1.
backendBackend to allocate the returned action Tensor through. Not owned; must outlive this object.
seedSeed for the internal deterministic LCG.
Exceptions
std::invalid_argumentif policy_network is null or action_dim <= 0.

Member Function Documentation

◆ action_dim()

int64_t pulsatrix::GFlowNetForwardPolicy::action_dim ( ) const
inline

Number of actions, as passed to the constructor.

◆ masked_probs()

std::vector< float > pulsatrix::GFlowNetForwardPolicy::masked_probs ( const Tensor &  observation,
const std::vector< bool > &  valid_actions 
)

The masked categorical distribution's probabilities, without sampling.

Parameters
observationObservation, shape (1, observation_dim).
valid_actionsSame mask contract as sample().
Returns
softmax(mask(policy_network(observation))), size action_dim().
Exceptions
Sameconditions as sample(), except no LCG draw is consumed and nothing is sampled – a training loop's two-pass backward step (mission_trajectory_balance_loss.md's Design section) needs to recompute the exact distribution sample() used, to derive the softmax/log gradient identity, without perturbing the LCG stream a second act() would.

◆ sample()

GFlowNetSampledAction pulsatrix::GFlowNetForwardPolicy::sample ( const Tensor &  observation,
const std::vector< bool > &  valid_actions 
)

Samples an action from the policy's distribution, restricted to valid_actions.

Parameters
observationObservation, shape (1, observation_dim) – whatever policy_network accepts.
valid_actionsBoolean mask, size action_dim(). At least one entry must be true.
Returns
The sampled action index (shape (1, 1) float Tensor) and the log-probability the masked distribution assigned to it.
Exceptions
std::invalid_argumentif policy_network's output is not (1, action_dim), if valid_actions.size() != action_dim(), or if every entry of valid_actions is false (no valid action to sample) – all external boundary.
Note
Exactly one LCG draw is consumed per call, unconditionally – same determinism convention as CategoricalPolicyAgent::act().

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