pulsatrix
Loading...
Searching...
No Matches
generator_population.hpp
Go to the documentation of this file.
1
22#pragma once
23
24#include <memory>
25#include <stdexcept>
26#include <vector>
27
28#include "pulsatrix/module.hpp"
29#include "pulsatrix/tensor.hpp"
30
31namespace pulsatrix {
32
35public:
37 explicit GeneratorPopulation(std::vector<std::unique_ptr<Module>> generators)
38 : generators_(std::move(generators)) {
39 if (generators_.empty()) {
40 throw std::invalid_argument("GeneratorPopulation: must have at least one generator");
41 }
42 for (const auto& g : generators_) {
43 if (g == nullptr) {
44 throw std::invalid_argument("GeneratorPopulation: generator entries must not be null");
45 }
46 }
47 }
48
49 [[nodiscard]] size_t size() const { return generators_.size(); }
50
52 [[nodiscard]] Module& generator(size_t index) { return *generators_.at(index); }
53 [[nodiscard]] const Module& generator(size_t index) const { return *generators_.at(index); }
54
62 void ReplaceGenerator(size_t index, std::unique_ptr<Module> new_generator) {
63 if (new_generator == nullptr) {
64 throw std::invalid_argument("GeneratorPopulation::ReplaceGenerator: new_generator must not be null");
65 }
66 generators_.at(index) = std::move(new_generator);
67 }
68
69private:
70 std::vector<std::unique_ptr<Module>> generators_;
71};
72
80inline Tensor SliceBatch(const Tensor& t, int64_t start, int64_t count, DeviceBackend* backend) {
81 if (t.shape().rank() < 1) {
82 throw std::invalid_argument("SliceBatch: t must have rank >= 1");
83 }
84 int64_t leading = t.shape().dim(0);
85 if (count <= 0 || start < 0 || start + count > leading) {
86 throw std::invalid_argument("SliceBatch: [start, start+count) must be a valid sub-range of t's leading dimension");
87 }
88
89 std::vector<int64_t> out_dims{count};
90 for (int64_t i = 1; i < t.shape().rank(); ++i) {
91 out_dims.push_back(t.shape().dim(i));
92 }
93 Shape out_shape(out_dims);
94
95 int64_t row_size = (leading > 0) ? (t.numel() / leading) : 0;
96 std::vector<float> values(static_cast<size_t>(row_size * count));
97 for (int64_t i = 0; i < row_size * count; ++i) {
98 values[static_cast<size_t>(i)] = t.data()[start * row_size + i];
99 }
100 return Tensor(out_shape, backend, values, t.device());
101}
102
109inline std::vector<float> FlattenParameters(Module& module) {
110 std::vector<float> flat;
111 for (const auto& p : module.parameters()) {
112 for (int64_t i = 0; i < p.value->numel(); ++i) {
113 flat.push_back(p.value->data()[i]);
114 }
115 }
116 return flat;
117}
118
125inline void RestoreParameters(Module& module, const std::vector<float>& flat) {
126 size_t offset = 0;
127 for (const auto& p : module.parameters()) {
128 int64_t n = p.value->numel();
129 for (int64_t i = 0; i < n; ++i) {
130 if (offset >= flat.size()) {
131 throw std::invalid_argument("RestoreParameters: flat has fewer entries than module.parameters() needs");
132 }
133 p.value->data()[i] = flat[offset];
134 ++offset;
135 }
136 }
137 if (offset != flat.size()) {
138 throw std::invalid_argument("RestoreParameters: flat has more entries than module.parameters() needs");
139 }
140}
141
149inline void ZeroModuleGradients(Module& module) {
150 for (const auto& p : module.parameters()) {
151 for (int64_t i = 0; i < p.grad->numel(); ++i) {
152 p.grad->data()[i] = 0.0f;
153 }
154 }
155}
156
166 const std::vector<Tensor>& noise_per_generator, DeviceBackend* backend) {
167 if (noise_per_generator.size() != population.size()) {
168 throw std::invalid_argument("GeneratePooledFakeSamples: noise_per_generator.size() must equal population.size()");
169 }
170 std::vector<Tensor> fakes;
171 fakes.reserve(population.size());
172 for (size_t i = 0; i < population.size(); ++i) {
173 fakes.push_back(population.generator(i).forward(noise_per_generator[i]));
174 }
175 return Tensor::Stack(fakes, backend);
176}
177
191inline void BackwardThroughPopulation(GeneratorPopulation& population, const Tensor& pooled_grad,
192 const std::vector<int64_t>& batch_sizes, DeviceBackend* backend) {
193 if (batch_sizes.size() != population.size()) {
194 throw std::invalid_argument("BackwardThroughPopulation: batch_sizes.size() must equal population.size()");
195 }
196 int64_t total = 0;
197 for (int64_t b : batch_sizes) {
198 total += b;
199 }
200 if (pooled_grad.shape().rank() < 1 || total != pooled_grad.shape().dim(0)) {
201 throw std::invalid_argument("BackwardThroughPopulation: batch_sizes must sum to pooled_grad's leading dimension");
202 }
203
204 int64_t offset = 0;
205 for (size_t i = 0; i < population.size(); ++i) {
206 int64_t count = batch_sizes[i];
207 Tensor grad_slice = SliceBatch(pooled_grad, offset, count, backend);
208 (void)population.generator(i).backward(grad_slice);
209 offset += count;
210 }
211}
212
213} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Owns N independently-parameterized generator Modules.
Definition generator_population.hpp:34
GeneratorPopulation(std::vector< std::unique_ptr< Module > > generators)
Definition generator_population.hpp:37
Module & generator(size_t index)
Definition generator_population.hpp:52
size_t size() const
Definition generator_population.hpp:49
const Module & generator(size_t index) const
Definition generator_population.hpp:53
void ReplaceGenerator(size_t index, std::unique_ptr< Module > new_generator)
Replaces population member index with new_generator – the concrete mechanism a future survivor-select...
Definition generator_population.hpp:62
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
Tensor forward(const Tensor &input)
Runs this module's forward computation.
Definition module.hpp:73
virtual Tensor backward(const Tensor &grad_output)=0
Computes the gradient w.r.t. this module's input, given the gradient w.r.t. its output....
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
int64_t dim(size_t index) const
Size of a single dimension.
Definition shape.hpp:93
int64_t rank() const
Number of dimensions. 0 for a scalar.
Definition shape.hpp:62
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
DeviceType device() const
Which device this tensor's buffer conceptually resides on.
Definition tensor.hpp:122
int64_t numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
const float * data() const
Raw buffer access. nullptr iff numel() == 0.
Definition tensor.hpp:167
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
Tensor SliceBatch(const Tensor &t, int64_t start, int64_t count, DeviceBackend *backend)
Extracts rows [start, start+count) along t's leading dimension into a new Tensor – the inverse of Ten...
Definition generator_population.hpp:80
std::vector< float > FlattenParameters(Module &module)
Flattens every parameter tensor module.parameters() reports (in that order) into one vector – the con...
Definition generator_population.hpp:109
void BackwardThroughPopulation(GeneratorPopulation &population, const Tensor &pooled_grad, const std::vector< int64_t > &batch_sizes, DeviceBackend *backend)
Given pooled_grad (the gradient w.r.t. the pooled fake batch GeneratePooledFakeSamples produced – e....
Definition generator_population.hpp:191
Tensor GeneratePooledFakeSamples(GeneratorPopulation &population, const std::vector< Tensor > &noise_per_generator, DeviceBackend *backend)
Runs every population member's generator forward on its own noise batch (noise_per_generator[i] for p...
Definition generator_population.hpp:165
void ZeroModuleGradients(Module &module)
Zeros every gradient tensor module.parameters() reports – a standalone alternative to calling some op...
Definition generator_population.hpp:149
void RestoreParameters(Module &module, const std::vector< float > &flat)
Overwrites every parameter tensor module.parameters() reports (in that order) from flat – the inverse...
Definition generator_population.hpp:125
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).