46 [[nodiscard]]
float forward(
float log_flow_i,
float sum_log_pf,
float log_flow_j,
float sum_log_pb,
47 float pair_weight_ratio);
50 [[nodiscard]]
float delta()
const;
70 float last_delta_ = 0.0f;
71 float last_pair_weight_ratio_ = 0.0f;
72 bool has_forwarded_ =
false;
74 void require_forwarded()
const;
One sub-trajectory pair's contribution to the SubTB(λ) loss: Δ(i,j) = log F(s_i) + Σ log P_F − log F(...
Definition subtb_loss.hpp:27
float grad_weight_for_log_pf_range() const
The weight w = pair_weight_ratio · (-2Δ(i,j)), applied uniformly to every edge t in [i,...
float forward(float log_flow_i, float sum_log_pf, float log_flow_j, float sum_log_pb, float pair_weight_ratio)
Computes this pair's Δ(i,j) and its weighted contribution to the total SubTB loss,...
float grad_log_flow_j() const
pair_weight_ratio · (-2Δ(i,j)). Only meaningful when s_j is non-terminal (a real F_θ network call) – ...
float grad_log_flow_i() const
pair_weight_ratio · 2Δ(i,j).
float delta() const
Δ(i,j) (unweighted) from the most recent forward() call.
Definition acquisition_functions.hpp:16