135 lines
4.6 KiB
C#

using System.Collections.Generic;
using UnityEngine;
namespace NanoBrain {
public class Population {
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;
}
/// <summary>
/// Generational update
/// </summary>
public virtual void Update() {
}
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-1f);
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;
}
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);
}
}
}