pulsatrix
Loading...
Searching...
No Matches
tanh_gaussian_policy.hpp
Go to the documentation of this file.
1
5#pragma once
6
9
10namespace pulsatrix {
11
29
38
72public:
82 static constexpr float kLogProbStabilizer = 1e-6f;
83
89
111 [[nodiscard]] TanhGaussianSample forward(const Tensor& mean, const Tensor& log_std, const Tensor& epsilon);
112
143 [[nodiscard]] TanhGaussianGrad backward(const Tensor& grad_action, const Tensor& grad_log_prob) const;
144
145private:
146 DeviceBackend* backend_;
147 // The cached action (not u) -- every backward term is expressible in a, std and epsilon,
148 // and caching a avoids recomputing tanh in backward().
149 Tensor last_action_;
150 Tensor last_std_;
151 Tensor last_epsilon_;
152 bool has_forwarded_ = false;
153};
154
155} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
SAC's reparameterized, tanh-squashed Gaussian policy sample (Haarnoja et al. 2018,...
Definition tanh_gaussian_policy.hpp:71
TanhGaussianSample forward(const Tensor &mean, const Tensor &log_std, const Tensor &epsilon)
Samples action = tanh(mean + exp(log_std)*epsilon) and its corrected log-density, caching what backwa...
static constexpr float kLogProbStabilizer
The fixed stabilizer added inside log(1 - action^2 + eps_stab).
Definition tanh_gaussian_policy.hpp:82
TanhGaussianPolicy(DeviceBackend *backend)
Constructs a tanh-Gaussian policy sampling step.
TanhGaussianGrad backward(const Tensor &grad_action, const Tensor &grad_log_prob) const
Gradients w.r.t. mean and log_std, combining the two gradients that arrive at this step's two outputs...
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
Definition acquisition_functions.hpp:16
The (grad_mean, grad_log_std) pair TanhGaussianPolicy::backward() produces – both (N,...
Definition tanh_gaussian_policy.hpp:34
Tensor grad_mean
Definition tanh_gaussian_policy.hpp:35
Tensor grad_log_std
Definition tanh_gaussian_policy.hpp:36
The (action, log_prob) pair TanhGaussianPolicy::forward() produces.
Definition tanh_gaussian_policy.hpp:25
Tensor log_prob
Definition tanh_gaussian_policy.hpp:27
Tensor action
Definition tanh_gaussian_policy.hpp:26
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).