pulsatrix
Loading...
Searching...
No Matches
trajectory_balance_loss.hpp
Go to the documentation of this file.
1
5#pragma once
6
7namespace pulsatrix {
8
23public:
25
34 [[nodiscard]] float forward(float sum_log_pf, float sum_log_pb, float log_reward, float log_z);
35
37 [[nodiscard]] float delta() const;
38
40 [[nodiscard]] float grad_log_z() const;
41
48 [[nodiscard]] float grad_weight_for_log_pf() const;
49
50private:
51 float last_delta_ = 0.0f;
52 bool has_forwarded_ = false;
53
56 void require_forwarded() const;
57};
58
59} // namespace pulsatrix
Δ(τ) = log Zθ + Σ log P_F(s_{t+1}|s_t) − log R(x) − Σ log P_B(s_t|s_{t+1}), loss = Δ(τ)².
Definition trajectory_balance_loss.hpp:22
float forward(float sum_log_pf, float sum_log_pb, float log_reward, float log_z)
Computes Δ(τ) and the loss Δ(τ)², caching Δ for the accessors below.
float grad_log_z() const
d(loss)/d(log Z) = 2Δ – feed directly into LearnableScalar::accumulate_grad.
float delta() const
Δ(τ) from the most recent forward() call.
float grad_weight_for_log_pf() const
The weight w = -2Δ such that the per-trajectory-step policy-network gradient is w · (p[k] − 1{k==a_t}...
Definition acquisition_functions.hpp:16