# IrisClassification.py

from operator import itemgetter
import math
import random

datafile = "iris.csv"
k = 3 # number of nearest neighbors
trainingQuota = 0.8 # relativ size of training set

def loadData(fileName):
    try:    
        fData = open(fileName, 'r')
    except:
        return []
    out = []
    for line in fData:
        line = line[:-1]  # remove \n
        if len(line) == 0:  # empty line
            continue
        li = [i for i in line.split(",")]
        out.append(li)
    fData.close()
    return out

def predict(sample):
    distances = []
    for i in range(len(trainingSet)):
        sum = 0
        for k in range(4):
            dk = float(trainingSet[i][k]) - float(sample[k])
            sum += dk * dk
        distance = math.sqrt(sum)    
        distances.append([i, distance])
    sorted_distances = sorted(distances, key = itemgetter(1))
    nearestSamples = []
    for i in range(k):
        nearestSamples.append(sorted_distances[i][0])
    votes = [0, 0, 0]
    for i in range(k):
        # get votes
        if trainingSet[nearestSamples[i]][4] == "Iris-setosa":
            votes[0] += 1
        elif trainingSet[nearestSamples[i]][4] == "Iris-versicolor":
            votes[1] += 1
        elif trainingSet[nearestSamples[i]][4] == "Iris-virginica":
            votes[2] += 1
    max_value = max(votes)
    max_index = votes.index(max_value)
    if max_index == 0:
        return "Iris-setosa"
    elif max_index == 1:
        return "Iris-versicolor"
    elif max_index == 2:
        return "Iris-virginica"

def splitSet(X, quota):
    randIndex = range(len(X))
    random.shuffle(randIndex)
    nbData = len(X)
    nbTraining = int(quota * nbData)
    training = [X[i] for i in randIndex[0:nbTraining]]  
    test = [X[i] for i in randIndex[nbTraining:nbData]] 
    return training, test
            
X = loadData(datafile)
trainingSet, testSet = splitSet(X, trainingQuota)
success = 0
for sample in testSet:
    p = predict(sample)
    if p == sample[4]:
        success += 1

print "Training set of size", len(trainingSet), ". Test set of size", len(testSet)
print "Result for", k, "nearest neighbors:",
print "Success", success, "out of", len(testSet), "samples -> ", \
    round(100 * success / len(testSet), 2), "percent" 
