update: changed NN and number of features
This commit is contained in:
+11
-6
@@ -5,11 +5,13 @@ 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, 32)
|
self.hidden1 = nn.Linear(inputSize, 32)
|
||||||
self.output = nn.Linear(32, 3)
|
self.hidden2 = nn.Linear(32, 16)
|
||||||
|
self.output = nn.Linear(16, 3)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
x = torch.relu(self.hidden(x))
|
x = torch.relu(self.hidden1(x))
|
||||||
|
x = torch.relu(self.hidden2(x))
|
||||||
x = self.output(x)
|
x = self.output(x)
|
||||||
return torch.argmax(x)
|
return torch.argmax(x)
|
||||||
|
|
||||||
@@ -28,9 +30,12 @@ class GeneticAlgorithm:
|
|||||||
def crossover(self, parent1: NeuralNetwork, parent2: NeuralNetwork):
|
def crossover(self, parent1: NeuralNetwork, parent2: NeuralNetwork):
|
||||||
child1 = NeuralNetwork(self.inputSize).to(self.device)
|
child1 = NeuralNetwork(self.inputSize).to(self.device)
|
||||||
child2 = NeuralNetwork(self.inputSize).to(self.device)
|
child2 = NeuralNetwork(self.inputSize).to(self.device)
|
||||||
point = len(child1.hidden.weight.data) // 2
|
point1 = len(child1.hidden1.weight.data) // 2
|
||||||
child1.hidden.weight.data = torch.cat((parent1.hidden.weight.data[:point], parent2.hidden.weight.data[point:]), dim=0)
|
point2 = len(child1.hidden2.weight.data) // 2
|
||||||
child2.hidden.weight.data = torch.cat((parent2.hidden.weight.data[:point], parent1.hidden.weight.data[point:]), dim=0)
|
child1.hidden1.weight.data = torch.cat((parent1.hidden1.weight.data[:point1], parent2.hidden1.weight.data[point1:]), dim=0)
|
||||||
|
child2.hidden1.weight.data = torch.cat((parent2.hidden1.weight.data[:point1], parent1.hidden1.weight.data[point1:]), dim=0)
|
||||||
|
child1.hidden2.weight.data = torch.cat((parent1.hidden2.weight.data[:point2], parent2.hidden2.weight.data[point2:]), dim=0)
|
||||||
|
child2.hidden2.weight.data = torch.cat((parent2.hidden2.weight.data[:point2], parent1.hidden2.weight.data[point2:]), dim=0)
|
||||||
child1.output.weight.data = parent1.output.weight.data.clone().detach()
|
child1.output.weight.data = parent1.output.weight.data.clone().detach()
|
||||||
child2.output.weight.data = parent2.output.weight.data.clone().detach()
|
child2.output.weight.data = parent2.output.weight.data.clone().detach()
|
||||||
return child1, child2
|
return child1, child2
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ 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
|
||||||
|
BARRIER_SPEED = 10
|
||||||
|
|
||||||
pygame.init()
|
pygame.init()
|
||||||
|
|
||||||
@@ -32,8 +33,6 @@ def init_game():
|
|||||||
grass = Background("images/Grass.jpg", 0, 0)
|
grass = Background("images/Grass.jpg", 0, 0)
|
||||||
grass.set_size(WIDTH, HEIGHT)
|
grass.set_size(WIDTH, HEIGHT)
|
||||||
|
|
||||||
BARRIER_SPEED = 1
|
|
||||||
|
|
||||||
gameobjects = []
|
gameobjects = []
|
||||||
|
|
||||||
barriers = []
|
barriers = []
|
||||||
@@ -65,12 +64,12 @@ def get_features(horse):
|
|||||||
last_barrier.rect.bottomright[0], last_barrier.rect.bottomright[1],
|
last_barrier.rect.bottomright[0], last_barrier.rect.bottomright[1],
|
||||||
horse.rect.topleft[0], horse.rect.topleft[1],
|
horse.rect.topleft[0], horse.rect.topleft[1],
|
||||||
horse.rect.bottomright[0], horse.rect.bottomright[1],
|
horse.rect.bottomright[0], horse.rect.bottomright[1],
|
||||||
HEIGHT, BARRIER_SPEED]
|
0, HEIGHT, BARRIER_SPEED]
|
||||||
return features
|
return features
|
||||||
|
|
||||||
genecticAlg = GeneticAlgorithm(POPULATION_SIZE, MUTATION_RATE, POPULATION_BEST, 10)
|
|
||||||
horses = [Horse("images/Horse_1.png", 50, HEIGHT/2 + 0 * i) for i in range(POPULATION_SIZE)]
|
horses = [Horse("images/Horse_1.png", 50, HEIGHT/2 + 0 * i) for i in range(POPULATION_SIZE)]
|
||||||
init_game()
|
init_game()
|
||||||
|
genecticAlg = GeneticAlgorithm(POPULATION_SIZE, MUTATION_RATE, POPULATION_BEST, len(get_features(horses[0])))
|
||||||
while True:
|
while True:
|
||||||
for event in pygame.event.get():
|
for event in pygame.event.get():
|
||||||
if event.type == pygame.QUIT: # Handle window close event
|
if event.type == pygame.QUIT: # Handle window close event
|
||||||
|
|||||||
Reference in New Issue
Block a user