import pandas as pd
import numpy as np
import math
#import matplotlib.pyplot as plt

##################################################        

# zachazeni s pandas dataframem

s = pd.Series([1,3,5,5,1])
s.value_counts()

df = pd.DataFrame({ 'A' : 1.,
                    'B' : pd.Timestamp('20130102'),
                    'C' : pd.Series(1,index=list(range(4)),dtype='float32'),
                    'D' : np.array([3] * 4,dtype='int32'),
                    'E' : pd.Categorical(["test","train","test","train"]),
                    'F' : 'foo' })

df.dtypes
df.head()
df.tail(2)

#sloupce
df.columns

#pristup ke sloupcum
df.E
df["E"]

#index (radky)
df.index

#selekce a projekce pres label

df.loc[[1,2],["A","B"]]   
df.loc[[1,2],:]   
df.loc[:,["A","B"]]   

df.loc[[True, True, False, True],:]

df.loc[lambda row: row.E == "train",:]

df.loc[df.E == "train",:]




columns = ["edible", "cap-shape", "cap-surface", "cap-color", "bruises?",
        "odor", "gill-attachment", "gill-spacing", "gill-size", "gill-color",
        "stalk-shape", "stalk-root", "stalk-surface-above-ring",
        "stalk-surface-below-ring", "stalk-color-above-ring",
        "stalk-color-below-ring", "veil-type", "veil-color", "ring-number",
        "ring-type", "spore-print-color", "population", "habitat"
        ]
dataset = pd.read_csv("mushroom.data",
                      names=columns, index_col=None)

X = dataset.drop("edible", axis=1)
y = dataset["edible"]



#columns = ["class",
#           "handicapped-infants",
#           "water-project-cost-sharing",
#           "adoption-of-the-budget-resolution",
#           "physician-fee-freeze",
#           "el-salvador-aid",
#           "religious-groups-in-schools",
#           "anti-satellite-test-ban",
#           "aid-to-nicaraguan-contras",
#           "mx-missile",
#           "immigration",
#           "synfuels-corporation-cutback",
#           "education-spending",
#           "superfund-right-to-sue",
#           "crime",
#           "duty-free-exports",
#           "export-administration-act-south-africa"]
#
#dataset = pd.read_csv("house-votes-84.data", names=columns, index_col=None,na_values="?")
#dataset = dataset.dropna().reset_index(drop=True)
#X = dataset.drop("class", axis=1)
#y = dataset["class"]


##################################################        

def H(y):
    total = len(y)    
    result = 0
    for count in y.value_counts():
        p = count/total
        result -= p * math.log2(p)
    return result


def reduction_of_impurity(y,y_splits,impurity=H):
    result = impurity(y)
    total = len(y)
    for y_split in y_splits:
        result -= len(y_split)/total * impurity(y_split)
    return result



##################################################        

class decision_tree:
    
    def __init__(self):
        self.root = None
        self.height = 0
    

    def stopping_criterion(self,X,y,lvl):
        #X   -- data v uzlu
        #y   -- jejich tridy
        #lvl -- v jake jsme hloubce
        
        #return(lvl==3)
        return(y.value_counts().size == 1)

        
    def splitting_criterion(self,y,y_splits):
        return reduction_of_impurity(y,y_splits)


    def fit(self,X,y):
        self.height = 0
        self.root = self.induce_tree(X,y,0)
        return self
   
    def predict(self,X):
        
        decisions = []
        for  _, row in X.iterrows():
        
            actualNode = self.root
            
            
            while (not actualNode.decisionNode):
                value = row[actualNode.attribute]
                try:
                    actualNode = actualNode.childs[value]
                except KeyError:
                    #trenovaci data tuto hodnotu
                    #(kombinaci hodnot)
                    #neobsahovala
                    #vracime "neznamo"
                    decisions.append("?")
                    break;
            else:        
                decisions.append(actualNode.decision)
        
        return pd.Series(decisions)
        
      
    def induce_tree(self,X,y,lvl=0):

        result = node()
        
        if self.stopping_criterion(X,y,lvl):
            
            result.decisionNode = True
            result.decision = y.value_counts().index[0]

        else:
                
            max_crit = -1
            max_attribute = None
            max_attr_values = None
            max_indices = None
            
            for col in X.columns:
                
                #prepare splits
                attr_values = X[col].unique()
                indices = [(X[col] == value) for value in attr_values]
                crit = self.splitting_criterion(y,[y[index] for index in indices])
                                
                if crit > max_crit:
                    max_crit = crit
                    max_attribute = col
                    max_indices = indices
                    max_attr_values = attr_values
        
            # OK, mame nejlepsi attribut na split
            # vygenerujeme potomky
            
            self.height = max(self.height,lvl)
            childs = {value : self.induce_tree(X[index] , y[index],lvl+1) 
            for value,index in zip(max_attr_values, max_indices) }
            result.decisionNode = False
            result.childs = childs
            result.attribute = max_attribute
            
        return result
                

               
        
        



##################################################        
        
class node:

    def __init__(self):

        # attribut, podle ktereho se rozhodujeme
        # ...
        self.attribute = None
        
        self.decisionNode = False
                
        # priznak, jestli je to rozhodovaci 
        # uzel (list)
        self.decision = None
        
        # seznam dvojic: hodnota atributu
        self.childs = []
        
    
##################################################        
DT = decision_tree()    
DT.fit(X,y)
DT.predict(X)

all(DT.predict(X)==y)

##################################################        
# k-fold cross-validation
# k = 10

from sklearn.model_selection import StratifiedKFold
from sklearn.metrics import classification_report
from sklearn.metrics import accuracy_score

k = 10
kf = StratifiedKFold(n_splits=k, shuffle=True)


acc = 0

for train_index, test_index in kf.split(X,y):
    X_train, X_test = X.iloc[train_index,:], X.iloc[test_index,:]
    y_train, y_test = y.iloc[train_index], y[test_index]


    DT.fit(X_train,y_train)    
    y_pred = DT.predict(X_test)
    print("height:",DT.height)

    acc += accuracy_score(y_test, y_pred)
    
print(acc/k)














    
