Δ(τ) = log Zθ + Σ log P_F(s_{t+1}|s_t) − log R(x) − Σ log P_B(s_t|s_{t+1}), loss = Δ(τ)².
More...
#include <trajectory_balance_loss.hpp>
|
| | TrajectoryBalanceLoss ()=default |
| |
| float | forward (float sum_log_pf, float sum_log_pb, float log_reward, float log_z) |
| | Computes Δ(τ) and the loss Δ(τ)², caching Δ for the accessors below.
|
| |
| float | delta () const |
| | Δ(τ) from the most recent forward() call.
|
| |
| float | grad_log_z () const |
| | d(loss)/d(log Z) = 2Δ – feed directly into LearnableScalar::accumulate_grad.
|
| |
| float | grad_weight_for_log_pf () const |
| | The weight w = -2Δ such that the per-trajectory-step policy-network gradient is w · (p[k] − 1{k==a_t}) – PolicyGradientLoss's exact softmax/log gradient identity shape, with this trajectory's Δ playing the role of the per-step return.
|
| |
Δ(τ) = log Zθ + Σ log P_F(s_{t+1}|s_t) − log R(x) − Σ log P_B(s_t|s_{t+1}), loss = Δ(τ)².
- Note
- Deliberately does not follow
MSELoss/PolicyGradientLoss's single-tensor backward() shape: Δ feeds two heterogeneous consumers (a bare scalar log Z, and a sequence of per-step policy-network logits from different forward-policy calls at different states) that no single gradient tensor covers. Instead exposes delta() (the one real error signal) plus two named derived accessors. See mission_trajectory_balance_loss.md's Design section for the full derivation.
-
Not a Module subclass, for
MSELoss's own reason: a loss is the seed point relevance propagation starts from, not something propagate_relevance is defined for.
◆ TrajectoryBalanceLoss()
| pulsatrix::TrajectoryBalanceLoss::TrajectoryBalanceLoss |
( |
| ) |
|
|
default |
◆ delta()
| float pulsatrix::TrajectoryBalanceLoss::delta |
( |
| ) |
const |
◆ forward()
| float pulsatrix::TrajectoryBalanceLoss::forward |
( |
float |
sum_log_pf, |
|
|
float |
sum_log_pb, |
|
|
float |
log_reward, |
|
|
float |
log_z |
|
) |
| |
Computes Δ(τ) and the loss Δ(τ)², caching Δ for the accessors below.
- Parameters
-
| sum_log_pf | Σ_t log P_F(s_{t+1}|s_t) over the whole trajectory. |
| sum_log_pb | Σ_t log P_B(s_t|s_{t+1}) over the whole trajectory. |
| log_reward | log R(x) at the trajectory's terminal state. |
| log_z | The current log Zθ estimate (LearnableScalar::value()). |
- Returns
- The scalar loss,
Δ(τ)².
◆ grad_log_z()
| float pulsatrix::TrajectoryBalanceLoss::grad_log_z |
( |
| ) |
const |
◆ grad_weight_for_log_pf()
| float pulsatrix::TrajectoryBalanceLoss::grad_weight_for_log_pf |
( |
| ) |
const |
The weight w = -2Δ such that the per-trajectory-step policy-network gradient is w · (p[k] − 1{k==a_t}) – PolicyGradientLoss's exact softmax/log gradient identity shape, with this trajectory's Δ playing the role of the per-step return.
The documentation for this class was generated from the following file: