pulsatrix
Loading...
Searching...
No Matches
egan_training.hpp
Go to the documentation of this file.
1
14#pragma once
15
16#include <limits>
17#include <memory>
18#include <random>
19#include <stdexcept>
20#include <vector>
21
26
27namespace pulsatrix {
28
42template <typename MakeGenerator, typename DOptimizerT>
43std::vector<float> RunEGANGeneration(GeneratorPopulation& population, Module& discriminator,
44 const Tensor& real_batch, const std::vector<Tensor>& noise_per_generator,
45 const std::vector<MutationObjective>& objectives, MakeGenerator make_generator,
46 float g_learning_rate, DOptimizerT& d_optimizer, DeviceBackend* backend,
47 float gamma = 0.05f) {
48 if (noise_per_generator.size() != population.size()) {
49 throw std::invalid_argument("RunEGANGeneration: noise_per_generator.size() must equal population.size()");
50 }
51 if (objectives.empty()) {
52 throw std::invalid_argument("RunEGANGeneration: objectives must not be empty");
53 }
54
55 std::vector<float> winning_fitness(population.size());
56 std::vector<int64_t> batch_sizes(population.size());
57
58 for (size_t i = 0; i < population.size(); ++i) {
59 std::vector<float> parent_weights = FlattenParameters(population.generator(i));
60 std::unique_ptr<Module> best_offspring;
61 float best_fitness = -std::numeric_limits<float>::infinity();
62
63 for (MutationObjective objective : objectives) {
64 std::unique_ptr<Module> offspring = make_generator();
65 RestoreParameters(*offspring, parent_weights);
66 SGDOptimizer g_optimizer(g_learning_rate);
67 float fitness =
68 RunMutationStep(*offspring, discriminator, noise_per_generator[i], objective, g_optimizer, backend, gamma);
69 d_optimizer.zero_grad(discriminator); // RunMutationStep's own documented contract
70
71 if (fitness > best_fitness) {
72 best_fitness = fitness;
73 best_offspring = std::move(offspring);
74 }
75 }
76
77 winning_fitness[i] = best_fitness;
78 batch_sizes[i] = noise_per_generator[i].shape().dim(0);
79 population.ReplaceGenerator(i, std::move(best_offspring));
80 }
81
82 // Discriminator step: real term + pooled-fake term from the now-mutated population.
83 BCEWithLogitsLoss bce(backend);
84 Tensor logits_real = discriminator.forward(real_batch);
85 Tensor real_labels(logits_real.shape(), backend);
86 real_labels.fill(1.0f);
87 (void)bce.forward(logits_real, real_labels);
88 (void)discriminator.backward(bce.backward());
89
90 Tensor pooled_fake = GeneratePooledFakeSamples(population, noise_per_generator, backend);
91 Tensor logits_fake = discriminator.forward(pooled_fake);
92 Tensor fake_labels(logits_fake.shape(), backend);
93 fake_labels.fill(0.0f);
94 (void)bce.forward(logits_fake, fake_labels);
95 Tensor grad_pooled_fake = discriminator.backward(bce.backward());
96 BackwardThroughPopulation(population, grad_pooled_fake, batch_sizes, backend);
97
98 d_optimizer.step(discriminator);
99 d_optimizer.zero_grad(discriminator);
100 // BackwardThroughPopulation just accumulated the discriminator step's own gradient
101 // contribution into every population member's parameter-gradient buffer (an unavoidable
102 // side effect of routing a gradient through backward(), the same contamination hazard
103 // BCEWithLogitsLoss's own @warning documents for the generator step) -- cleared here so
104 // next generation's own mutation attempts start from a clean buffer, not a mix of this
105 // generation's D-step contamination and their own fresh accumulation.
106 for (size_t i = 0; i < population.size(); ++i) {
107 ZeroModuleGradients(population.generator(i));
108 }
109
110 return winning_fitness;
111}
112
119template <typename MakeGenerator, typename DOptimizerT, typename RNG>
120void RunEGANTraining(GeneratorPopulation& population, Module& discriminator, const Tensor& real_batch,
121 int num_generations, int64_t noise_dim, int64_t per_generator_batch,
122 const std::vector<MutationObjective>& objectives, MakeGenerator make_generator,
123 float g_learning_rate, DOptimizerT& d_optimizer, DeviceBackend* backend, RNG& rng,
124 float gamma = 0.05f) {
125 if (num_generations <= 0) {
126 throw std::invalid_argument("RunEGANTraining: num_generations must be positive");
127 }
128
129 std::normal_distribution<float> noise_dist(0.0f, 1.0f);
130 for (int gen = 0; gen < num_generations; ++gen) {
131 std::vector<Tensor> noise_per_generator;
132 noise_per_generator.reserve(population.size());
133 for (size_t i = 0; i < population.size(); ++i) {
134 std::vector<float> values(static_cast<size_t>(per_generator_batch * noise_dim));
135 for (auto& v : values) {
136 v = noise_dist(rng);
137 }
138 noise_per_generator.emplace_back(Shape({per_generator_batch, noise_dim}), backend, values);
139 }
140 (void)RunEGANGeneration(population, discriminator, real_batch, noise_per_generator, objectives,
141 make_generator, g_learning_rate, d_optimizer, backend, gamma);
142 }
143}
144
145} // namespace pulsatrix
Binary cross-entropy on raw logits – combined sigmoid + BCE, numerically stable.
loss = mean( max(x,0) - x*y + log(1 + exp(-|x|)) ), over all N*k elements of a (N,...
Definition bce_with_logits_loss.hpp:50
Tensor backward() const
Gradient w.r.t. the logits: grad[i] = (sigmoid(x[i]) - y[i]) / numel.
float forward(const Tensor &logits, const Tensor &target)
Computes the loss value and caches logits/target for backward().
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
Module & generator(size_t index)
Definition generator_population.hpp:52
size_t size() const
Definition generator_population.hpp:49
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....
param -= learning_rate * grad, per parameter, for every parameter a Module exposes.
Definition sgd_optimizer.hpp:12
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Tensor & fill(float value)
Sets every element to value. Safe no-op on a zero-element tensor.
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
E-GAN's own three named "mutation" objectives (Wang et al. 2019, "Evolutionary Generative Adve...
E-GAN's own population-of-generators infrastructure (Wang et al. 2019, "Evolutionary Generativ...
Definition acquisition_functions.hpp:16
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 RunEGANTraining(GeneratorPopulation &population, Module &discriminator, const Tensor &real_batch, int num_generations, int64_t noise_dim, int64_t per_generator_batch, const std::vector< MutationObjective > &objectives, MakeGenerator make_generator, float g_learning_rate, DOptimizerT &d_optimizer, DeviceBackend *backend, RNG &rng, float gamma=0.05f)
RNG-driven wrapper: draws fresh standard-normal noise for every population member every generation,...
Definition egan_training.hpp:120
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
std::vector< float > RunEGANGeneration(GeneratorPopulation &population, Module &discriminator, const Tensor &real_batch, const std::vector< Tensor > &noise_per_generator, const std::vector< MutationObjective > &objectives, MakeGenerator make_generator, float g_learning_rate, DOptimizerT &d_optimizer, DeviceBackend *backend, float gamma=0.05f)
Runs one E-GAN generation: for every population member, attempts every objective in objectives (each ...
Definition egan_training.hpp:43
void ZeroModuleGradients(Module &module)
Zeros every gradient tensor module.parameters() reports – a standalone alternative to calling some op...
Definition generator_population.hpp:149
MutationObjective
E-GAN's three named mutation objectives.
Definition egan_mutation.hpp:36
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
float RunMutationStep(Module &offspring, Module &discriminator, const Tensor &noise, MutationObjective objective, OptimizerT &g_optimizer, DeviceBackend *backend, float gamma=0.05f)
Runs one E-GAN mutation training step: trains offspring (an already-independent Module instance – typ...
Definition egan_mutation.hpp:182
Stochastic gradient descent – operates uniformly across any Module's parameters().