# the soma function for Orcs vs Uruk-Hais
def predict(speed, straight_legged, threshold, w1, w2):
    x = 0
    x = speed * w1 + straight_legged * w2 + threshold
    if x > 0:
        return 1
    else:
        return 0

#import the JMSSNeural functions
from JMSSNeural import *

#this is your training data!
#[speed, straight-leggedness, isorc]
#Extension: read this from your .csv file instead of typing it all in again!
training = []

file = open("ORCS_URUK.csv", "r")
n = 0
for rows in file:
    if n > 0:
        coloumns = rows.split(",")
        line = []
        line.append(float(coloumns[0]))
        line.append(float(coloumns[1]))
        line.append(float(coloumns[2]))
        training.append(line)
    else:
        n += 1





"""
training = [
[0.1, 0.1, 1],
[0.2, 0.1, 1],
[0.3, 0.2, 1],
[0.1, 0.4, 1],
[0.6, 0.1, 1],
[0.5, 0.3, 1],
[0.2, 0.5, 1],
[0.3, 0.6, 1],
[0.2, 0.7, 1],
[0.1, 0.7, 1],
[0.9, 0.2, 0],
[0.7, 0.4, 0],
[0.6, 0.5, 0],
[0.9, 0.5, 0],
[0.8, 0.6, 0],
[0.7, 0.7, 0],
[0.9, 0.8, 0],
[0.8, 0.8, 0],
[0.5, 0.8, 0],
[1, 1, 0], 
]"""


#define your learning rate and number of epochs
#this is up to you!
learning_rate = 0.1
epochs = 50

# this will conduct the training over the specified number of epochs and
# using the learning rate
weights = train(training, learning_rate, epochs)

# show me the answers for the weights
# weights[0] = threshold
# weights[1] = w1
# weights[2] = w2
print("The trained weights are: ", weights)


# now use the weights (outputted from the previous run)
# to PREDICT whether the creature is an orc or not given:
#   a speed and straight-leggedness value
print(predict(speed = 0.9, \
           straight_legged = 0.5, \
           threshold = 0.2, \
           w1 = -0.29000000000000015, \
           w2 = -0.09000000000000001))
print(predict(speed = 0.28, \
           straight_legged = 0.1, \
           threshold = 0.2, \
           w1 = -0.29000000000000015, \
           w2 = -0.09000000000000001))
print(predict(speed = 0.35, \
           straight_legged = 0.3, \
           threshold = 0.2, \
           w1 = -0.29000000000000015, \
           w2 = -0.09000000000000001))
print(predict(speed = 0.6, \
           straight_legged = 0.8, \
           threshold = 0.2, \
           w1 = -0.29000000000000015, \
           w2 = -0.09000000000000001))
print(predict(speed = 0.8, \
           straight_legged = 0.7, \
           threshold = 0.2, \
           w1 = -0.29000000000000015, \
           w2 = -0.09000000000000001))
print(predict(speed = 1, \
           straight_legged = 1, \
           threshold = 0.2, \
           w1 = -0.29000000000000015, \
           w2 = -0.09000000000000001))
file.close()
