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

‘Δ(s,s’) = log F(s) + log P_F(s'|s) − log F(s') − log P_B(s|s'),loss = Δ(s,s')²-- the per-*transition* credit-assignment alternative toTrajectoryBalanceLoss`'s per-trajectory constraint. More...

#include <detailed_balance_loss.hpp>

Public Member Functions

 DetailedBalanceLoss ()=default
 
float forward (float log_flow_s, float log_pf, float log_flow_s_next, float log_pb)
 Computes ‘Δ(s,s’)and the lossΔ(s,s')², cachingΔfor the accessors below. @param log_flow_slog F_θ(s). @param log_pflog P_F(s'|s)(orlog P_F(stop|s)for the exit transition). @param log_flow_s_nextlog F_θ(s')for a non-terminals', orlog R(x)for the exit transition (s' = x, the trajectory's terminal state). @param log_pblog P_B(s|s')for a real transition, or0.0for the exit transition. @return The scalar loss,Δ(s,s')²`.
 
float delta () const
 ‘Δ(s,s’)` from the most recent forward() call.
 
float grad_log_flow_s () const
 d(loss)/d(log F(s)) = 2Δ.
 
float grad_log_flow_s_next () const
 ‘d(loss)/d(log F(s’)) = -2Δ‘. Only meaningful when s’ is non-terminal (a real F_θ network call) – callers must not apply this to the exit transition's fixed log R(x) endpoint, which has no network to backprop into.
 
float grad_weight_for_log_pf () const
 The weight w = -2Δ such that this transition's policy-network gradient is w · (p[k] − 1{k==a}) – TrajectoryBalanceLoss/PolicyGradientLoss's exact softmax/log gradient identity shape.
 

Detailed Description

‘Δ(s,s’) = log F(s) + log P_F(s'|s) − log F(s') − log P_B(s|s'),loss = Δ(s,s')²-- the per-*transition* credit-assignment alternative toTrajectoryBalanceLoss`'s per-trajectory constraint.

Note
The terminal/exit transition is not a special case of this class: passing log_f_s_next = log R(x) (the reward, not a learned flow) and log_pb = 0.0 (log(1), the trivial backward policy from a unique virtual sink state) reduces the general formula to Δ_exit = log F(s_n) + log P_F(stop|s_n) − log R(x) exactly. See mission_detailed_balance_loss.md's Design section.
Structurally identical math to TrajectoryBalanceLoss (a + b − c − d, squared) at a different granularity – implemented as its own class rather than sharing a base, matching this codebase's own precedent (MSELoss/CrossEntropyLoss/ BCEWithLogitsLoss/KLDivergenceLoss are all independent despite similar shapes).
Not a Module subclass, for MSELoss's own reason.

Constructor & Destructor Documentation

◆ DetailedBalanceLoss()

pulsatrix::DetailedBalanceLoss::DetailedBalanceLoss ( )
default

Member Function Documentation

◆ delta()

float pulsatrix::DetailedBalanceLoss::delta ( ) const

‘Δ(s,s’)` from the most recent forward() call.

◆ forward()

float pulsatrix::DetailedBalanceLoss::forward ( float  log_flow_s,
float  log_pf,
float  log_flow_s_next,
float  log_pb 
)

Computes ‘Δ(s,s’)and the lossΔ(s,s')², cachingΔfor the accessors below. @param log_flow_slog F_θ(s). @param log_pflog P_F(s'|s)(orlog P_F(stop|s)for the exit transition). @param log_flow_s_nextlog F_θ(s')for a non-terminals', orlog R(x)for the exit transition (s' = x, the trajectory's terminal state). @param log_pblog P_B(s|s')for a real transition, or0.0for the exit transition. @return The scalar loss,Δ(s,s')²`.

◆ grad_log_flow_s()

float pulsatrix::DetailedBalanceLoss::grad_log_flow_s ( ) const

d(loss)/d(log F(s)) = 2Δ.

◆ grad_log_flow_s_next()

float pulsatrix::DetailedBalanceLoss::grad_log_flow_s_next ( ) const

‘d(loss)/d(log F(s’)) = -2Δ‘. Only meaningful when s’ is non-terminal (a real F_θ network call) – callers must not apply this to the exit transition's fixed log R(x) endpoint, which has no network to backprop into.

◆ grad_weight_for_log_pf()

float pulsatrix::DetailedBalanceLoss::grad_weight_for_log_pf ( ) const

The weight w = -2Δ such that this transition's policy-network gradient is w · (p[k] − 1{k==a}) – TrajectoryBalanceLoss/PolicyGradientLoss's exact softmax/log gradient identity shape.


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