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>
|
| | 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.
|
| |
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.
◆ 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_network | Network mapping a (1, observation_dim) observation to (1, action_dim) raw logits. Not owned; must outlive this object. |
| action_dim | Number of actions. Must be >= 1. |
| backend | Backend to allocate the returned action Tensor through. Not owned; must outlive this object. |
| seed | Seed for the internal deterministic LCG. |
- Exceptions
-
| std::invalid_argument | if policy_network is null or action_dim <= 0. |
◆ 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
-
| observation | Observation, shape (1, observation_dim). |
| valid_actions | Same mask contract as sample(). |
- Returns
softmax(mask(policy_network(observation))), size action_dim().
- Exceptions
-
| Same | conditions 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
-
| observation | Observation, shape (1, observation_dim) – whatever policy_network accepts. |
| valid_actions | Boolean 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_argument | if 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: