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

Δ(τ) = 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>

Public Member Functions

 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.
 

Detailed Description

Δ(τ) = 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.

Constructor & Destructor Documentation

◆ TrajectoryBalanceLoss()

pulsatrix::TrajectoryBalanceLoss::TrajectoryBalanceLoss ( )
default

Member Function Documentation

◆ delta()

float pulsatrix::TrajectoryBalanceLoss::delta ( ) const

Δ(τ) from the most recent forward() call.

◆ 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_rewardlog R(x) at the trajectory's terminal state.
log_zThe current log Zθ estimate (LearnableScalar::value()).
Returns
The scalar loss, Δ(τ)².

◆ grad_log_z()

float pulsatrix::TrajectoryBalanceLoss::grad_log_z ( ) const

d(loss)/d(log Z) = 2Δ – feed directly into LearnableScalar::accumulate_grad.

◆ 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: