import numpy as np
import pandas as pd

#Print you can execute arbitrary python code
train = pd.read_csv("../input/train.csv", dtype={"Age": np.float64}, )
test = pd.read_csv("../input/test.csv", dtype={"Age": np.float64}, )

# splitPoints = ['gender', 'class_1', 'class_2','hasFam', 'child', 'hasCabin']
splitPoints = ['gender', 'class_1', 'class_2', 'child',  'hasFam', 'hasCabin']

# validation break
vnum = 0

# preprocessing
def preprocess(dataSet):
    dataSet['gender'] = dataSet['Sex'].apply(lambda sex: sex == 'female')
    dataSet['class_1'] = dataSet['Pclass'].apply(lambda pclass: pclass == 1)
    dataSet['class_2'] = dataSet['Pclass'].apply(lambda pclass: pclass == 2)
    dataSet['hasCabin'] = dataSet['Cabin'].apply(lambda cabin: cabin is not None)
    dataSet['child'] = dataSet['Age'].apply(lambda age: age < 8)
    dataSet['hasFam'] = dataSet['Parch'].apply(lambda par : par > 0)
    
    # drop data we don't use
    del dataSet['Name']
    del dataSet['Sex']
    del dataSet['Age']
    del dataSet['Ticket']
    del dataSet['Pclass']
    del dataSet['Parch']
    del dataSet['Fare']
    del dataSet['Cabin']
    del dataSet['Embarked']

def postprocess(dataSet):
    if 'Surv' in dataSet:
        dataSet['Survived'] = dataSet['Surv']
    for field in dataSet:
        if field != 'PassengerId' and field != 'Survived':
            del dataSet[field]

preprocess(test)
preprocess(train)

class DecisionTree(object):
    def __init__(self, dataSet, splitPoints):
        self.splitPointList = splitPoints
        splitPoint, truthy, falsy, rate = bestSplit(dataSet, splitPoints)
        # print (truthy.info(), falsy.info())
        self.rate = rate
        if splitPoint is None:
            # leaf node
            self.splitPoint = None
            self.isSurvivor = rate > 0.5
            self.n = len(dataSet)
        else:
            if (len(truthy) == 0 or len(falsy) == 0):
                print ("::" + splitPoint)
            self.nums = len(truthy), len(falsy)
            newPoints = [point for point in splitPoints if point != splitPoint]
            self.splitPoint = splitPoint
            self.truthy = DecisionTree(truthy, newPoints)
            self.falsy = DecisionTree(falsy, newPoints)
    def __str__(self):
        return self.toStr(0)
    def toStr(self, depth):
        if self.splitPoint is None:
            return ("\t" * depth) + str(self.isSurvivor) + "[" + str(self.rate)  + ", " + str(self.n) + "]" + "\n"
        else:
            return ("\t" * depth) + "(" + self.splitPoint + str(self.nums) + str(self.rate) + "\n" + self.truthy.toStr(depth + 1) +  self.falsy.toStr(depth + 1) + ("\t" * depth) + ")" + "\n"
    def evaluate(self, record):
        if self.splitPoint is None:
            return 1 if self.isSurvivor else 0
        else:
            if record[self.splitPoint]:
                return 1 if self.truthy.evaluate(record) else 0
            else:
                return 1 if self.falsy.evaluate(record) else 0
    def evaluate_all(self, dataSet):
        dataSet['Surv'] = dataSet.apply(self.evaluate, 1)

# work out how good each split is, return the best
#
# returns: (bestSplit, truthy, falsy, rate)
def bestSplit(dataSet, points, verbose = False):
    surv = dataSet['Survived'] == 1
    if len(points) == 0:
        # end of recursion
        return (None, None, None, len(dataSet[surv]) / len(dataSet))
    if len(dataSet) == 0:
        # out of data, error
        assert False
    size = len(dataSet)
    bestRate = 0.0 # worst possible
    bestPoint = None
    bestRatePair = (0.0, 0.0)
    for point in points: # evaluate this split
        # filters
        t = dataSet[point]
        f = dataSet[point] == False
        # counts
        trueCount = len(dataSet[t])
        falseCount = len(dataSet[f])
        trueSurvCount = len(dataSet[t & surv])
        falseSurvCount = len(dataSet[f & surv])
        if trueCount == 0 or falseCount == 0:
            # empty side, zero information gain whatsoever
            continue
        # rates, want choice which maximizes the dependency
        trueRate = trueSurvCount / trueCount # rate of survival for true passengers
        falseRate = falseSurvCount / falseCount # rate of survival for false passengers
        rate = abs(trueRate - falseRate)
        #rate = max(rate, 1 - rate)
        if verbose:
            print(point, rate, trueRate, falseRate)
        # update best rate
        if rate >= bestRate:
            bestRate = rate
            bestPoint = point
            bestRatePair = (trueRate, falseRate)
    if bestPoint is None:
        return (None, None, None, len(dataSet[surv]) / len(dataSet))
    return (bestPoint, dataSet[dataSet[bestPoint]], dataSet[dataSet[bestPoint] == False], bestRatePair)

if vnum > 0:
    validation = train[:vnum]
    train = train[vnum:]
else:
    validation = train[:]
tree = DecisionTree(train, splitPoints)

tree.evaluate_all(train)
tree.evaluate_all(validation)

def evaluateTree(dataSet):
    tp = len(dataSet[(dataSet['Survived'] == 1) & (dataSet['Surv'] == 1)])
    tn = len(dataSet[(dataSet['Survived'] == 0) & (dataSet['Surv'] == 0)])
    fp = len(dataSet[(dataSet['Survived'] == 0) & (dataSet['Surv'] == 1)])
    fn = len(dataSet[(dataSet['Survived'] == 1) & (dataSet['Surv'] == 0)])
    
    print ('.', '\t', 'P', '\t', 'N')
    print ('T', '\t', tp,'\t', tn, '\t', tp / (tp  + tn))
    print ('F', '\t', fp,'\t', fn, '\t', fp / (fp + fn))
    print ('.', '\t', tp / (tp + fp), '\t', tn / (tn + fn))
    print ((tp + tn) / (tp + tn + fp + fn))
    print (len(dataSet[dataSet['Survived'] == dataSet['Surv']]) / len(dataSet))

print (evaluateTree(train))
print (evaluateTree(validation))

tree.evaluate_all(test)
postprocess(test)

#Any files you save will be available in the output tab below
train.to_csv('out_train.csv', index = False)
test.to_csv('out_test.csv', index = False)