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 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.Copy(); for (int memberIx = 1; memberIx < members.Count; memberIx++) { Member member = members[memberIx]; result.ProcessWeightsFrom(member.brain, Sum); } result.ProcessWeights(w => w / members.Count); return result; } #endregion Saving #region New Generation public void NewGeneration() { foreach (Member member in this.members) { member.initialized = false; } } protected List SelectElite(float percentage) { return SelectElite((int)(members.Count * percentage)); } protected List SelectElite(int count) { List 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 GenerateMutants(List elite, float percentage) { return GenerateMutants(elite, (int)(members.Count * percentage)); } protected List GenerateMutants(List elite, int count) { List 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 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; } private static float Sum(float a, float b) { return a + b; } protected List GenerateRandom() { return GenerateRandom(int.MaxValue); } protected List GenerateRandom(int count) { List 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() { } } }