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

Samples an action from softmax(policy_network(observation)) – the discrete-action stochastic policy every policy-gradient method in this phase (REINFORCE, A2C, PPO) is built on. More...

#include <categorical_policy_agent.hpp>

Inheritance diagram for pulsatrix::CategoricalPolicyAgent:
Collaboration diagram for pulsatrix::CategoricalPolicyAgent:

Public Member Functions

 CategoricalPolicyAgent (Module *policy_network, int64_t action_dim, DeviceBackend *backend, uint32_t seed=42)
 Constructs a categorical policy over a logit-producing network.
 
Tensor act (const Tensor &observation) override
 Samples an action from the policy's categorical distribution, advancing the LCG.
 
Tensor act_greedy (const Tensor &observation)
 The deterministic (evaluation-time) policy: argmax_a logits[a], no sampling.
 
float log_prob () const
 The log-probability the policy assigned to the most recent act() sample.
 
int64_t action_dim () const
 Number of discrete actions, as passed to the constructor.
 
- Public Member Functions inherited from pulsatrix::Agent
virtual ~Agent ()=default
 

Detailed Description

Samples an action from softmax(policy_network(observation)) – the discrete-action stochastic policy every policy-gradient method in this phase (REINFORCE, A2C, PPO) is built on.

Note
The policy network outputs raw logits, not probabilities. The numerically stable softmax (subtract the row max before exponentiating) is computed here, fused into the agent, for exactly the reason CrossEntropyLoss fuses softmax with NLL rather than composing a SoftmaxModule with a separate log: the log-probability this agent must report is logit - max - log(sum exp(logit - max)), which never materializes a probability that could underflow to zero and then be logged.
Takes a Module* rather than defining a policy-network type of its own, and does not own it – identical convention to DQNAgent's q_network. Must outlive this agent.
log_prob() is an extra method beyond the Agent interface, not an extension of it. Agent::act() returns only an action Tensor, but policy-gradient methods need the log-probability the policy assigned to the action it just sampled (RolloutBuffer stores it). Resolved the same way DQNAgent already added act_greedy()/set_epsilon() beyond the interface: a non-interface method reading cached forward-pass state, the same convention every Module already uses for its own last_*_ cache.
Sampling randomness comes from an internal deterministic LCG seeded at construction – never \<random\>, whose engine outputs beyond mt19937 are implementation-defined. Byte-for-byte the generator and constants CartPoleEnv, ReplayBuffer and DQNAgent use, which is the only thing that makes a stochastic policy hand-traceable in a test.

Constructor & Destructor Documentation

◆ CategoricalPolicyAgent()

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

Constructs a 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 agent.
action_dimNumber of discrete actions. Must be >= 1.
backendBackend to allocate the returned action Tensor through. Not owned; must outlive this agent.
seedSeed for the internal deterministic LCG.
Exceptions
std::invalid_argumentif policy_network is null or action_dim <= 0 – external boundaries, same classification as DQNAgent's identical checks.

Member Function Documentation

◆ act()

Tensor pulsatrix::CategoricalPolicyAgent::act ( const Tensor &  observation)
overridevirtual

Samples an action from the policy's categorical distribution, advancing the LCG.

Parameters
observationObservation, shape (1, observation_dim) – whatever policy_network accepts.
Returns
The sampled action index encoded as a (1, 1) float Tensor, identical to the encoding Environment::step() expects for a discrete environment.
Exceptions
std::invalid_argumentif policy_network's output is not (1, action_dim) – external boundary: the network and action_dim are two independent constructor arguments a caller can genuinely mismatch, and an unchecked scan would then read past the output row.
Note
Exactly one LCG draw is consumed per call, unconditionally: u uniform in [0, 1), then action is the smallest index whose cumulative probability reaches u (inverse-CDF sampling – a single cumulative-sum scan, not a rejection loop, so the number of LCG steps is content-independent and the stream stays easy to reason about in a determinism test).
Caches the sampled action's log-probability for log_prob(), taken straight from the stable log-softmax rather than re-derived as log(p[action]) – one rounding path, not two possibly inconsistent ones.
Host boundary (GPU-native-kernels Mission 7): sampling (the softmax and the inverse-CDF scan against the LCG draw) is host logic. The network may run on any device; its (1, action_dim) logits are copied to the host once per call, and the action is returned through this agent's own backend. act_greedy() likewise.

Implements pulsatrix::Agent.

◆ act_greedy()

Tensor pulsatrix::CategoricalPolicyAgent::act_greedy ( const Tensor &  observation)

The deterministic (evaluation-time) policy: argmax_a logits[a], no sampling.

Parameters
observationObservation, shape (1, observation_dim).
Returns
The greedy action index as a (1, 1) float Tensor.
Exceptions
std::invalid_argumentif policy_network's output is not (1, action_dim).
Note
argmax over the logits is argmax over the softmax probabilities – softmax is strictly monotone – so no exponentiation is needed here at all.
Consumes no randomness, advances no state, and deliberately does not update log_prob(): an evaluation rollout interleaved with training can therefore neither perturb the training run's action sequence nor corrupt the log-probability the training loop is about to store for the last genuinely sampled action. That is what makes this a separate method rather than "call act() and ignore the randomness", mirroring DQNAgent::act_greedy exactly.
Not const despite its name's suggestion of purity: Module::forward() is non-const (every module caches forward state for its own backward()), so running the network necessarily mutates it – identical reasoning to DQNAgent::act_greedy.

◆ action_dim()

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

Number of discrete actions, as passed to the constructor.

◆ log_prob()

float pulsatrix::CategoricalPolicyAgent::log_prob ( ) const

The log-probability the policy assigned to the most recent act() sample.

Returns
log_softmax[action] for the action act() last returned.
Exceptions
std::logic_errorif act() has never been called – there is no meaningful "probability of no action", and returning 0.0f (a probability of 1) would be a silently plausible lie that a training loop would happily backpropagate.
Note
Unaffected by act_greedy(), which samples nothing.

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