179 lines
5.8 KiB
C#

using System.Collections.Generic;
using UnityEngine;
namespace NanoBrain {
public class Population {
# region Members
public class Member {
public Member(Cluster cluster) {
this.brain = cluster;
this.performance = 0;
this.initialized = true;
}
public Cluster brain;
public float performance;
public bool initialized;
}
public List<Member> members = new();
public Member AddMember(Cluster brain) {
Member newMember = new(brain);
this.members.Add(newMember);
return newMember;
}
#endregion Members
#region Saving
public enum SaveMethod {
Average,
Mean,
First,
All
}
public void Save(string path, SaveMethod saveMethod = SaveMethod.Average) {
Cluster cluster = null;
switch (saveMethod) {
case SaveMethod.First:
Member member = this.members[0];
if (member == null)
return;
cluster = member.brain;
break;
case SaveMethod.Average:
cluster = CalculateAverageCluster();
break;
}
cluster?.Export(path);
}
protected Cluster CalculateAverageCluster() {
if (members.Count <= 0)
return null;
Cluster result = members[0].brain.Clone();
for (int memberIx = 1; memberIx < members.Count; memberIx++) {
Member member = members[memberIx];
result.ProcessWeightsFrom(member.brain, Average);
}
return result;
}
#endregion Saving
#region New Generation
public void NewGeneration() {
foreach (Member member in this.members) {
member.initialized = false;
}
}
protected List<Member> SelectElite(float percentage) {
return SelectElite((int)(members.Count * percentage));
}
protected List<Member> SelectElite(int count) {
List<Member> selectedMembers = new();
SortMembers();
for (int i = 0; i < count; i++) {
Member member = this.members[i];
if (!member.initialized) {
Debug.Log($"Selected {i}: {member.performance}");
member.initialized = true;
selectedMembers.Add(member);
// LogTrainableWeights(member.brain);
}
}
return selectedMembers;
}
private void SortMembers() {
members.Sort((a, b) => a.performance.CompareTo(b.performance));
}
protected List<Member> GenerateMutants(List<Member> elite, float percentage) {
return GenerateMutants(elite, (int)(members.Count * percentage));
}
protected List<Member> GenerateMutants(List<Member> elite, int count) {
List<Member> selectedMembers = new();
int i = 0;
int n = 0;
while (n < count && i < this.members.Count) {
Member member = this.members[i];
if (member.initialized == false) {
GenerateMutant(member, elite);
member.initialized = true;
n++;
selectedMembers.Add(member);
}
i++;
}
return selectedMembers;
}
protected void GenerateMutant(Member member, List<Member> ants) {
System.Random randomGenerator = new();
int n = ants.Count;
int parent1 = randomGenerator.Next(0, n);
int parent2 = randomGenerator.Next(0, n);
Debug.Log($"Mutate from {members[parent1].performance} and {members[parent2].performance}");
member.brain.CopyWeightsFrom(ants[parent1].brain);
member.brain.ProcessWeightsFrom(ants[parent2].brain, Average);
// LogTrainableWeights(member.brain);
member.brain.GaussianAdditiveMutation(1e-2f);
// LogTrainableWeights(member.brain);
}
private static float Average(float a, float b) {
return (a + b) / 2;
}
protected List<Member> GenerateRandom() {
return GenerateRandom(int.MaxValue);
}
protected List<Member> GenerateRandom(int count) {
List<Member> selectedMembers = new();
int i = 0;
int n = 0;
while (n < count && i < this.members.Count) {
Member member = this.members[i];
if (member.initialized == false) {
Debug.Log($"Randomized {i}: {member.performance}");
member.brain.InitializeRandom();
member.initialized = true;
selectedMembers.Add(member);
n++;
}
i++;
}
return selectedMembers;
}
#endregion New Generation
private void LogTrainableWeights(Cluster cluster) {
string s = "";
foreach (Nucleus nucleus in cluster.nuclei) {
if (nucleus is Neuron neuron) {
foreach (Synapse synapse in neuron.synapses) {
if (synapse.trainable) {
s += synapse.weight + " ";
}
}
}
}
Debug.Log(s);
}
public virtual void Update() { }
}
}