pulsatrix
Loading...
Searching...
No Matches
egan_mutation.hpp
Go to the documentation of this file.
1
21#pragma once
22
23#include <cmath>
24#include <limits>
25#include <stdexcept>
26#include <vector>
27
29#include "pulsatrix/module.hpp"
31#include "pulsatrix/tensor.hpp"
32
33namespace pulsatrix {
34
37
46public:
47 explicit MutationLoss(DeviceBackend* backend) : backend_(backend), bce_(backend), mse_(backend) {}
48
50 float Forward(MutationObjective objective, const Tensor& logits) {
51 objective_ = objective;
52 Tensor ones(logits.shape(), backend_);
53 ones.fill(1.0f);
54 switch (objective) {
56 return bce_.forward(logits, ones);
58 return mse_.forward(logits, ones);
60 Tensor zeros(logits.shape(), backend_);
61 zeros.fill(0.0f);
62 return -bce_.forward(logits, zeros);
63 }
64 }
65 throw std::invalid_argument("MutationLoss::Forward: unknown MutationObjective");
66 }
67
73 [[nodiscard]] Tensor Backward() const {
74 switch (objective_) {
76 return bce_.backward();
78 return mse_.backward();
80 return Negate(bce_.backward());
81 }
82 throw std::invalid_argument("MutationLoss::Backward: unknown MutationObjective");
83 }
84
85private:
86 [[nodiscard]] Tensor Negate(const Tensor& t) const {
87 std::vector<float> values(static_cast<size_t>(t.numel()));
88 for (int64_t i = 0; i < t.numel(); ++i) {
89 values[static_cast<size_t>(i)] = -t.data()[i];
90 }
91 return Tensor(t.shape(), backend_, values, t.device());
92 }
93
94 DeviceBackend* backend_;
95 BCEWithLogitsLoss bce_;
96 MSELoss mse_;
98};
99
105inline float QualityFitness(const Tensor& logits) {
106 double total = 0.0;
107 for (int64_t i = 0; i < logits.numel(); ++i) {
108 total += 1.0 / (1.0 + std::exp(-static_cast<double>(logits.data()[i])));
109 }
110 return static_cast<float>(total / static_cast<double>(logits.numel()));
111}
112
129inline float DiversityFitness(Module& discriminator, const Tensor& fake, DeviceBackend* backend) {
130 auto params = discriminator.parameters();
131 if (params.empty()) {
132 throw std::invalid_argument("DiversityFitness: discriminator must have at least one parameter");
133 }
134
135 BCEWithLogitsLoss bce(backend);
136 Tensor logits = discriminator.forward(fake);
137 Tensor zeros(logits.shape(), backend);
138 zeros.fill(0.0f);
139 (void)bce.forward(logits, zeros);
140 (void)discriminator.backward(bce.backward());
141
142 double sum_squares = 0.0;
143 for (const auto& p : params) {
144 for (int64_t i = 0; i < p.grad->numel(); ++i) {
145 double g = p.grad->data()[i];
146 sum_squares += g * g;
147 }
148 }
149 double norm = std::sqrt(sum_squares);
150 // An exactly-zero gradient (e.g. against a freshly-constructed, all-zero-weight
151 // discriminator) maps to +infinity fitness -- treated as "maximally diverse" by explicit
152 // convention rather than left as an undefined log(0), a documented edge case, not a
153 // silent bug.
154 if (norm <= 0.0) {
155 return std::numeric_limits<float>::infinity();
156 }
157 return static_cast<float>(-std::log(norm));
158}
159
166inline float CombinedFitness(float quality, float diversity, float gamma = 0.05f) {
167 return quality + gamma * diversity;
168}
169
181template <typename OptimizerT>
182float RunMutationStep(Module& offspring, Module& discriminator, const Tensor& noise, MutationObjective objective,
183 OptimizerT& g_optimizer, DeviceBackend* backend, float gamma = 0.05f) {
184 Tensor fake = offspring.forward(noise);
185 Tensor logits = discriminator.forward(fake);
186 MutationLoss loss(backend);
187 (void)loss.Forward(objective, logits);
188 Tensor grad_fake = discriminator.backward(loss.Backward());
189 (void)offspring.backward(grad_fake);
190 g_optimizer.step(offspring);
191 g_optimizer.zero_grad(offspring);
192
193 Tensor post_fake = offspring.forward(noise);
194 Tensor post_logits = discriminator.forward(post_fake);
195 float quality = QualityFitness(post_logits);
196 float diversity = DiversityFitness(discriminator, post_fake, backend);
197 return CombinedFitness(quality, diversity, gamma);
198}
199
200} // 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
float forward(const Tensor &prediction, const Tensor &target)
Computes the loss value and caches prediction/target for backward().
Tensor backward() const
Computes the gradient w.r.t. the prediction: (2/n) * (prediction - target).
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
virtual std::vector< ParamRef > parameters()
This module's trainable parameters and their gradients, for an optimizer to update uniformly across m...
Definition module.hpp:180
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....
Computes one of E-GAN's three named mutation objectives against a discriminator's own raw logit outpu...
Definition egan_mutation.hpp:45
float Forward(MutationObjective objective, const Tensor &logits)
Computes the chosen objective's scalar value against logits.
Definition egan_mutation.hpp:50
MutationLoss(DeviceBackend *backend)
Definition egan_mutation.hpp:47
Tensor Backward() const
Gradient w.r.t. the logits passed to the most recent Forward() call.
Definition egan_mutation.hpp:73
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
Tensor & fill(float value)
Sets every element to value. Safe no-op on a zero-element tensor.
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.
Mean squared error loss.
Definition acquisition_functions.hpp:16
float DiversityFitness(Module &discriminator, const Tensor &fake, DeviceBackend *backend)
E-GAN's own diversity fitness Fd = -log(||grad||): the negative log of the L2 norm of the discriminat...
Definition egan_mutation.hpp:129
float QualityFitness(const Tensor &logits)
E-GAN's own quality fitness Fq: mean sigmoid(D(fake)) over the batch – how convincingly "real" the di...
Definition egan_mutation.hpp:105
MutationObjective
E-GAN's three named mutation objectives.
Definition egan_mutation.hpp:36
float CombinedFitness(float quality, float diversity, float gamma=0.05f)
Combined E-GAN fitness: Fq + gamma*Fd (Wang et al. 2019's own weighted combination)....
Definition egan_mutation.hpp:166
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
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).