pulsatrix
Loading...
Searching...
No Matches
subtb_loss.hpp
Go to the documentation of this file.
1
6#pragma once
7
8namespace pulsatrix {
9
27class SubTBLoss {
28public:
29 SubTBLoss() = default;
30
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);
48
50 [[nodiscard]] float delta() const;
51
53 [[nodiscard]] float grad_log_flow_i() const;
54
59 [[nodiscard]] float grad_log_flow_j() const;
60
67 [[nodiscard]] float grad_weight_for_log_pf_range() const;
68
69private:
70 float last_delta_ = 0.0f;
71 float last_pair_weight_ratio_ = 0.0f;
72 bool has_forwarded_ = false;
73
74 void require_forwarded() const;
75};
76
77} // namespace pulsatrix
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