pulsatrix
Loading...
Searching...
No Matches
detailed_balance_loss.hpp
Go to the documentation of this file.
1
5#pragma once
6
7namespace pulsatrix {
8
26public:
28
38 [[nodiscard]] float forward(float log_flow_s, float log_pf, float log_flow_s_next, float log_pb);
39
41 [[nodiscard]] float delta() const;
42
44 [[nodiscard]] float grad_log_flow_s() const;
45
49 [[nodiscard]] float grad_log_flow_s_next() const;
50
56 [[nodiscard]] float grad_weight_for_log_pf() const;
57
58private:
59 float last_delta_ = 0.0f;
60 bool has_forwarded_ = false;
61
62 void require_forwarded() const;
63};
64
65} // namespace 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...
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)....
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