|
pulsatrix
|
‘Δ(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. | |
‘Δ(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.
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. 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). MSELoss's own reason.
|
default |
| float pulsatrix::DetailedBalanceLoss::delta | ( | ) | const |
‘Δ(s,s’)` from the most recent forward() call.
| 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')²`.
| float pulsatrix::DetailedBalanceLoss::grad_log_flow_s | ( | ) | const |
d(loss)/d(log F(s)) = 2Δ.
| 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.
| 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.