|
pulsatrix
|
One full sampled trajectory: every state the forward policy acted from, the action taken at each, and the trajectory-level quantities Trajectory Balance (and Detailed Balance/SubTB, in later missions) need. More...
#include <gflownet_trajectory.hpp>
Public Attributes | |
| std::vector< Tensor > | states |
| The state the policy acted from at each decision point, in order. | |
| std::vector< int64_t > | actions |
| The action index sampled at each corresponding entry of states. | |
| float | sum_log_pf = 0.0f |
Σ_t log P_F(s_{t+1}|s_t) over the whole trajectory, including the final stop decision. | |
| float | sum_log_pb = 0.0f |
Σ_t log P_B(s_t|s_{t+1}) over every real move (the stop action does not move the state, so it contributes no P_B term). | |
| float | terminal_reward = 0.0f |
R(x) at the trajectory's terminal state. | |
One full sampled trajectory: every state the forward policy acted from, the action taken at each, and the trajectory-level quantities Trajectory Balance (and Detailed Balance/SubTB, in later missions) need.
states/actions deliberately are NOT enough on their own to recompute sum_log_pf/sum_log_pb – they exist so a training loop's second pass can re-forward() the policy network at each state (refreshing its single-most-recent-call cache) and apply the now-known per-step gradient. See mission_trajectory_balance_loss.md's Design section ("Two-pass rollout"). | std::vector<int64_t> pulsatrix::GFlowNetTrajectory::actions |
The action index sampled at each corresponding entry of states.
| std::vector<Tensor> pulsatrix::GFlowNetTrajectory::states |
The state the policy acted from at each decision point, in order.
| float pulsatrix::GFlowNetTrajectory::sum_log_pb = 0.0f |
Σ_t log P_B(s_t|s_{t+1}) over every real move (the stop action does not move the state, so it contributes no P_B term).
| float pulsatrix::GFlowNetTrajectory::sum_log_pf = 0.0f |
Σ_t log P_F(s_{t+1}|s_t) over the whole trajectory, including the final stop decision.
| float pulsatrix::GFlowNetTrajectory::terminal_reward = 0.0f |
R(x) at the trajectory's terminal state.