38 [[nodiscard]]
float forward(
float log_flow_s,
float log_pf,
float log_flow_s_next,
float log_pb);
41 [[nodiscard]]
float delta()
const;
59 float last_delta_ = 0.0f;
60 bool has_forwarded_ =
false;
62 void require_forwarded()
const;
‘Δ(s,s’) = log F(s) + log P_F(s'|s) − log F(s') − log P_B(s|s'),loss = Δ(s,s')²-- the per-*transition...
Definition detailed_balance_loss.hpp:25
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)....
DetailedBalanceLoss()=default
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) – cal...
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}) – Traj...
Definition acquisition_functions.hpp:16