pulsatrix
Loading...
Searching...
No Matches
neat_genome.hpp
Go to the documentation of this file.
1
16#pragma once
17
18#include <algorithm>
19#include <map>
20#include <random>
21#include <stdexcept>
22#include <utility>
23#include <vector>
24
25namespace pulsatrix {
26
28struct NodeGene {
29 enum class Type { Input, Bias, Output, Hidden };
30 int id;
32};
33
41 double weight;
42 bool enabled;
44};
45
56public:
57 explicit InnovationTracker(int next_node_id) : next_node_id_(next_node_id) {}
58
62 int GetConnectionInnovation(int in_node, int out_node) {
63 auto key = std::make_pair(in_node, out_node);
64 auto it = connection_innovations_.find(key);
65 if (it != connection_innovations_.end()) {
66 return it->second;
67 }
68 int innovation = next_innovation_++;
69 connection_innovations_[key] = innovation;
70 return innovation;
71 }
72
76 int GetNodeIdForSplit(int split_connection_innovation) {
77 auto it = node_ids_by_split_.find(split_connection_innovation);
78 if (it != node_ids_by_split_.end()) {
79 return it->second;
80 }
81 int node_id = next_node_id_++;
82 node_ids_by_split_[split_connection_innovation] = node_id;
83 return node_id;
84 }
85
86private:
87 std::map<std::pair<int, int>, int> connection_innovations_;
88 std::map<int, int> node_ids_by_split_;
89 int next_innovation_ = 0;
90 int next_node_id_;
91};
92
95public:
103 NEATGenome(int num_inputs, int num_outputs, bool has_bias, InnovationTracker& tracker) {
104 if (num_inputs <= 0 || num_outputs <= 0) {
105 throw std::invalid_argument("NEATGenome: num_inputs and num_outputs must be positive");
106 }
107 int next_id = 0;
108 std::vector<int> input_ids;
109 for (int i = 0; i < num_inputs; ++i) {
110 nodes_.push_back(NodeGene{next_id, NodeGene::Type::Input});
111 input_ids.push_back(next_id);
112 ++next_id;
113 }
114 if (has_bias) {
115 nodes_.push_back(NodeGene{next_id, NodeGene::Type::Bias});
116 input_ids.push_back(next_id);
117 ++next_id;
118 }
119 std::vector<int> output_ids;
120 for (int i = 0; i < num_outputs; ++i) {
121 nodes_.push_back(NodeGene{next_id, NodeGene::Type::Output});
122 output_ids.push_back(next_id);
123 ++next_id;
124 }
125 for (int in_id : input_ids) {
126 for (int out_id : output_ids) {
127 int innovation = tracker.GetConnectionInnovation(in_id, out_id);
128 connections_.push_back(ConnectionGene{in_id, out_id, 0.0, true, innovation});
129 }
130 }
131 }
132
133 [[nodiscard]] const std::vector<NodeGene>& nodes() const { return nodes_; }
134 [[nodiscard]] const std::vector<ConnectionGene>& connections() const { return connections_; }
135
144 void SetConnectionWeight(int innovation, double weight) {
145 auto it = std::find_if(connections_.begin(), connections_.end(),
146 [&](const ConnectionGene& c) { return c.innovation == innovation; });
147 if (it == connections_.end()) {
148 throw std::invalid_argument("NEATGenome::SetConnectionWeight: no connection with that innovation");
149 }
150 it->weight = weight;
151 }
152
160 void AddConnectionBetween(int in_node, int out_node, double weight, InnovationTracker& tracker) {
161 if (!HasNode(in_node) || !HasNode(out_node)) {
162 throw std::invalid_argument("NEATGenome::AddConnectionBetween: node not found in this genome");
163 }
164 for (const auto& c : connections_) {
165 if (c.in_node == in_node && c.out_node == out_node) {
166 throw std::invalid_argument("NEATGenome::AddConnectionBetween: connection already exists");
167 }
168 }
169 int innovation = tracker.GetConnectionInnovation(in_node, out_node);
170 connections_.push_back(ConnectionGene{in_node, out_node, weight, true, innovation});
171 }
172
180 template <typename RNG>
181 bool AddConnection(InnovationTracker& tracker, RNG& rng) {
182 std::vector<std::pair<int, int>> candidates;
183 for (const auto& a : nodes_) {
184 if (a.type == NodeGene::Type::Output) {
185 continue; // outputs never originate a connection in a feedforward network
186 }
187 for (const auto& b : nodes_) {
188 if (b.type == NodeGene::Type::Input || b.type == NodeGene::Type::Bias || a.id == b.id) {
189 continue;
190 }
191 if (ConnectionExists(a.id, b.id)) {
192 continue;
193 }
194 if (CanReach(b.id, a.id)) {
195 continue; // adding a.id -> b.id would close a cycle
196 }
197 candidates.emplace_back(a.id, b.id);
198 }
199 }
200 if (candidates.empty()) {
201 return false;
202 }
203 std::uniform_int_distribution<size_t> pick(0, candidates.size() - 1);
204 std::normal_distribution<double> weight_dist(0.0, 1.0);
205 auto [in_node, out_node] = candidates[pick(rng)];
206 AddConnectionBetween(in_node, out_node, weight_dist(rng), tracker);
207 return true;
208 }
209
220 void AddNodeSplitting(int connection_innovation, InnovationTracker& tracker) {
221 auto it = std::find_if(connections_.begin(), connections_.end(), [&](const ConnectionGene& c) {
222 return c.innovation == connection_innovation && c.enabled;
223 });
224 if (it == connections_.end()) {
225 throw std::invalid_argument("NEATGenome::AddNodeSplitting: no enabled connection with that innovation");
226 }
227 it->enabled = false;
228 int in_node = it->in_node;
229 int out_node = it->out_node;
230 double original_weight = it->weight;
231
232 int new_node_id = tracker.GetNodeIdForSplit(connection_innovation);
233 nodes_.push_back(NodeGene{new_node_id, NodeGene::Type::Hidden});
234
235 int innovation_in = tracker.GetConnectionInnovation(in_node, new_node_id);
236 connections_.push_back(ConnectionGene{in_node, new_node_id, 1.0, true, innovation_in});
237 int innovation_out = tracker.GetConnectionInnovation(new_node_id, out_node);
238 connections_.push_back(ConnectionGene{new_node_id, out_node, original_weight, true, innovation_out});
239 }
240
245 template <typename RNG>
246 bool AddNode(InnovationTracker& tracker, RNG& rng) {
247 std::vector<int> enabled_innovations;
248 for (const auto& c : connections_) {
249 if (c.enabled) {
250 enabled_innovations.push_back(c.innovation);
251 }
252 }
253 if (enabled_innovations.empty()) {
254 return false;
255 }
256 std::uniform_int_distribution<size_t> pick(0, enabled_innovations.size() - 1);
257 AddNodeSplitting(enabled_innovations[pick(rng)], tracker);
258 return true;
259 }
260
270 template <typename RNG>
271 void MutateWeights(double sigma, double mutation_probability, RNG& rng) {
272 if (sigma < 0.0) {
273 throw std::invalid_argument("NEATGenome::MutateWeights: sigma must be non-negative");
274 }
275 if (mutation_probability < 0.0 || mutation_probability > 1.0) {
276 throw std::invalid_argument("NEATGenome::MutateWeights: mutation_probability must be in [0, 1]");
277 }
278 std::bernoulli_distribution mask(mutation_probability);
279 std::normal_distribution<double> noise(0.0, sigma);
280 for (auto& c : connections_) {
281 if (c.enabled && mask(rng)) {
282 c.weight += noise(rng);
283 }
284 }
285 }
286
287private:
288 [[nodiscard]] bool HasNode(int id) const {
289 return std::any_of(nodes_.begin(), nodes_.end(), [&](const NodeGene& n) { return n.id == id; });
290 }
291
292 [[nodiscard]] bool ConnectionExists(int in_node, int out_node) const {
293 return std::any_of(connections_.begin(), connections_.end(), [&](const ConnectionGene& c) {
294 return c.in_node == in_node && c.out_node == out_node;
295 });
296 }
297
300 [[nodiscard]] bool CanReach(int from_node, int out_node) const {
301 std::vector<int> stack{from_node};
302 std::vector<int> visited;
303 while (!stack.empty()) {
304 int current = stack.back();
305 stack.pop_back();
306 if (current == out_node) {
307 return true;
308 }
309 if (std::find(visited.begin(), visited.end(), current) != visited.end()) {
310 continue;
311 }
312 visited.push_back(current);
313 for (const auto& c : connections_) {
314 if (c.enabled && c.in_node == current) {
315 stack.push_back(c.out_node);
316 }
317 }
318 }
319 return false;
320 }
321
322 std::vector<NodeGene> nodes_;
323 std::vector<ConnectionGene> connections_;
324};
325
326} // namespace pulsatrix
The global historical-marking registry: the same structural mutation (an identical new connection,...
Definition neat_genome.hpp:55
InnovationTracker(int next_node_id)
Definition neat_genome.hpp:57
int GetNodeIdForSplit(int split_connection_innovation)
The new node ID created by splitting the connection identified by split_connection_innovation – reuse...
Definition neat_genome.hpp:76
int GetConnectionInnovation(int in_node, int out_node)
The innovation number for a new (in_node, out_node) connection – reused if this exact connection has ...
Definition neat_genome.hpp:62
A NEAT genome: its node and connection genes, growable via structural mutation.
Definition neat_genome.hpp:94
const std::vector< ConnectionGene > & connections() const
Definition neat_genome.hpp:134
void MutateWeights(double sigma, double mutation_probability, RNG &rng)
Non-structural mutation: perturbs every enabled connection's weight independently with probability mu...
Definition neat_genome.hpp:271
void AddConnectionBetween(int in_node, int out_node, double weight, InnovationTracker &tracker)
Pure core: adds a new enabled connection gene (in_node, out_node, weight), assigning its innovation n...
Definition neat_genome.hpp:160
bool AddConnection(InnovationTracker &tracker, RNG &rng)
RNG-driven wrapper: proposes a random, currently-nonexistent, feedforward-safe (cannot create a cycle...
Definition neat_genome.hpp:181
NEATGenome(int num_inputs, int num_outputs, bool has_bias, InnovationTracker &tracker)
Constructs the minimal starting topology: num_inputs input nodes (+1 bias node if has_bias),...
Definition neat_genome.hpp:103
void SetConnectionWeight(int innovation, double weight)
Directly sets an existing connection's weight by innovation number – needed infrastructure found nece...
Definition neat_genome.hpp:144
const std::vector< NodeGene > & nodes() const
Definition neat_genome.hpp:133
void AddNodeSplitting(int connection_innovation, InnovationTracker &tracker)
Pure core: splits the enabled connection with the given innovation number – disables it (kept,...
Definition neat_genome.hpp:220
bool AddNode(InnovationTracker &tracker, RNG &rng)
RNG-driven wrapper: splits a uniformly-randomly chosen enabled connection. A no-op (returns false) if...
Definition neat_genome.hpp:246
Definition acquisition_functions.hpp:16
One connection in a NEAT genome: an edge between two node IDs, its weight, whether it is currently ac...
Definition neat_genome.hpp:38
int out_node
Definition neat_genome.hpp:40
int in_node
Definition neat_genome.hpp:39
double weight
Definition neat_genome.hpp:41
bool enabled
Definition neat_genome.hpp:42
int innovation
Definition neat_genome.hpp:43
One node in a NEAT genome's topology.
Definition neat_genome.hpp:28
Type type
Definition neat_genome.hpp:31
Type
Definition neat_genome.hpp:29
int id
Definition neat_genome.hpp:30