add: horses can learn now
This commit is contained in:
+7
-13
@@ -5,8 +5,8 @@ import numpy as np
|
|||||||
class NeuralNetwork(nn.Module):
|
class NeuralNetwork(nn.Module):
|
||||||
def __init__(self, inputSize):
|
def __init__(self, inputSize):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden = nn.Linear(inputSize, 10)
|
self.hidden = nn.Linear(inputSize, 32)
|
||||||
self.output = nn.Linear(10, 3)
|
self.output = nn.Linear(32, 3)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
x = torch.relu(self.hidden(x))
|
x = torch.relu(self.hidden(x))
|
||||||
@@ -19,16 +19,16 @@ class GeneticAlgorithm:
|
|||||||
self.mutationRate = mutationRate
|
self.mutationRate = mutationRate
|
||||||
self.percentageBest = percentageBest
|
self.percentageBest = percentageBest
|
||||||
self.inputSize = inputSize
|
self.inputSize = inputSize
|
||||||
|
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
self.initialize_population()
|
self.initialize_population()
|
||||||
|
|
||||||
def initialize_population(self):
|
def initialize_population(self):
|
||||||
self.population = [NeuralNetwork(self.inputSize) for _ in range(self.populationSize)]
|
self.population = [NeuralNetwork(self.inputSize).to(self.device) for _ in range(self.populationSize)]
|
||||||
|
|
||||||
def crossover(self, parent1: NeuralNetwork, parent2: NeuralNetwork):
|
def crossover(self, parent1: NeuralNetwork, parent2: NeuralNetwork):
|
||||||
child1 = NeuralNetwork()
|
child1 = NeuralNetwork(self.inputSize).to(self.device)
|
||||||
child2 = NeuralNetwork()
|
child2 = NeuralNetwork(self.inputSize).to(self.device)
|
||||||
point = len(child1.hidden.weight.data) // 2
|
point = len(child1.hidden.weight.data) // 2
|
||||||
print([x for x in child1.parameters()])
|
|
||||||
child1.hidden.weight.data = torch.cat((parent1.hidden.weight.data[:point], parent2.hidden.weight.data[point:]), dim=0)
|
child1.hidden.weight.data = torch.cat((parent1.hidden.weight.data[:point], parent2.hidden.weight.data[point:]), dim=0)
|
||||||
child2.hidden.weight.data = torch.cat((parent2.hidden.weight.data[:point], parent1.hidden.weight.data[point:]), dim=0)
|
child2.hidden.weight.data = torch.cat((parent2.hidden.weight.data[:point], parent1.hidden.weight.data[point:]), dim=0)
|
||||||
child1.output.weight.data = parent1.output.weight.data.clone().detach()
|
child1.output.weight.data = parent1.output.weight.data.clone().detach()
|
||||||
@@ -55,11 +55,5 @@ class GeneticAlgorithm:
|
|||||||
while len(self.population) > self.populationSize: self.population.pop()
|
while len(self.population) > self.populationSize: self.population.pop()
|
||||||
|
|
||||||
def predict(self, data: list, i):
|
def predict(self, data: list, i):
|
||||||
data = torch.tensor(data, requires_grad=False).float()
|
data = torch.tensor(data, requires_grad=False).float().to(self.device)
|
||||||
return self.population[i](data)
|
return self.population[i](data)
|
||||||
|
|
||||||
def predict_all(self, data: list):
|
|
||||||
result = []
|
|
||||||
for i in range(self.populationSize):
|
|
||||||
result.append(self.population[i](data[i]))
|
|
||||||
return result
|
|
||||||
@@ -3,7 +3,7 @@ import sys
|
|||||||
from gameobjects import *
|
from gameobjects import *
|
||||||
from genetic_alg import GeneticAlgorithm
|
from genetic_alg import GeneticAlgorithm
|
||||||
|
|
||||||
POPULATION_SIZE = 5
|
POPULATION_SIZE = 50
|
||||||
MUTATION_RATE = 0.5
|
MUTATION_RATE = 0.5
|
||||||
POPULATION_BEST = 0.2
|
POPULATION_BEST = 0.2
|
||||||
FRAME_RATE = 160 #TODO: FIX THE INCORRECT FRAME RATE CORRELATION
|
FRAME_RATE = 160 #TODO: FIX THE INCORRECT FRAME RATE CORRELATION
|
||||||
@@ -45,6 +45,7 @@ def init_game():
|
|||||||
horse.set_vacceleration(0)
|
horse.set_vacceleration(0)
|
||||||
horse.set_vspeed(0)
|
horse.set_vspeed(0)
|
||||||
horse.frame_counter = 0
|
horse.frame_counter = 0
|
||||||
|
horse.fitness = 0
|
||||||
spawner = Spawner("images/Barrier.png")
|
spawner = Spawner("images/Barrier.png")
|
||||||
new_barrier = spawner.spawn()
|
new_barrier = spawner.spawn()
|
||||||
barriers.append(new_barrier)
|
barriers.append(new_barrier)
|
||||||
@@ -124,6 +125,7 @@ while True:
|
|||||||
|
|
||||||
if all(horse.stopped for horse in horses):
|
if all(horse.stopped for horse in horses):
|
||||||
UI.iteration_num += 1
|
UI.iteration_num += 1
|
||||||
|
genecticAlg.learn([x.fitness for x in horses])
|
||||||
init_game()
|
init_game()
|
||||||
|
|
||||||
pygame.display.flip() # Update the display
|
pygame.display.flip() # Update the display
|
||||||
|
|||||||
Reference in New Issue
Block a user