pulsatrix
Loading...
Searching...
No Matches
gflownet_forward_policy.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <cstdint>
9#include <vector>
10
12#include "pulsatrix/module.hpp"
13#include "pulsatrix/tensor.hpp"
14
15namespace pulsatrix {
16
23
48public:
59 GFlowNetForwardPolicy(Module* policy_network, int64_t action_dim, DeviceBackend* backend, uint32_t seed = 42);
60
74 [[nodiscard]] GFlowNetSampledAction sample(const Tensor& observation, const std::vector<bool>& valid_actions);
75
87 [[nodiscard]] std::vector<float> masked_probs(const Tensor& observation, const std::vector<bool>& valid_actions);
88
90 [[nodiscard]] int64_t action_dim() const { return action_dim_; }
91
92private:
94 [[nodiscard]] float next_unit();
95
96 Module* policy_network_;
97 int64_t action_dim_;
98 DeviceBackend* backend_;
99 uint32_t lcg_state_;
100};
101
102} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Samples from softmax(mask(policy_network(observation))) – a categorical policy restricted to a caller...
Definition gflownet_forward_policy.hpp:47
GFlowNetForwardPolicy(Module *policy_network, int64_t action_dim, DeviceBackend *backend, uint32_t seed=42)
Constructs a masked categorical policy over a logit-producing network.
std::vector< float > masked_probs(const Tensor &observation, const std::vector< bool > &valid_actions)
The masked categorical distribution's probabilities, without sampling.
GFlowNetSampledAction sample(const Tensor &observation, const std::vector< bool > &valid_actions)
Samples an action from the policy's distribution, restricted to valid_actions.
int64_t action_dim() const
Number of actions, as passed to the constructor.
Definition gflownet_forward_policy.hpp:90
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
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.
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
One sample from GFlowNetForwardPolicy::sample: the chosen action and the log-probability the masked d...
Definition gflownet_forward_policy.hpp:19
Tensor action
Definition gflownet_forward_policy.hpp:20
float log_prob
Definition gflownet_forward_policy.hpp:21
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).