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>
|
| | 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.
|
| |
| virtual | ~Agent ()=default |
| |
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.
◆ 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_network | Network mapping a (1, observation_dim) observation to (1, action_dim) raw logits. Not owned; must outlive this agent. |
| action_dim | Number of discrete actions. Must be >= 1. |
| backend | Backend to allocate the returned action Tensor through. Not owned; must outlive this agent. |
| seed | Seed for the internal deterministic LCG. |
- Exceptions
-
| std::invalid_argument | if policy_network is null or action_dim <= 0 – external boundaries, same classification as DQNAgent's identical checks. |
◆ act()
| Tensor pulsatrix::CategoricalPolicyAgent::act |
( |
const Tensor & |
observation | ) |
|
|
overridevirtual |
Samples an action from the policy's categorical distribution, advancing the LCG.
- Parameters
-
| observation | Observation, 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_argument | if 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
-
| observation | Observation, shape (1, observation_dim). |
- Returns
- The greedy action index as a (1, 1) float Tensor.
- Exceptions
-
| std::invalid_argument | if 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_error | if 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: