import re
import operator

import numpy as np
import pandas as pd
import keras.backend as K
import tensorflow as tf

from gensim.models import KeyedVectors
from sklearn.model_selection import train_test_split
from keras.callbacks import EarlyStopping
from sklearn.metrics import accuracy_score
from sklearn.metrics import f1_score
from sklearn.linear_model import LinearRegression
from keras.layers import BatchNormalization
from keras.models import Sequential
from keras.models import Model
from keras.layers import Dense
from keras.layers import Dropout
from keras.layers import Embedding
from keras.layers import Flatten
from keras.layers import Conv1D
from keras.layers import MaxPooling1D
from keras.layers import Input
from keras.layers import Conv2D
from keras.layers import MaxPool2D
from keras.layers import Reshape
from keras.layers import Concatenate
from keras.layers import SpatialDropout1D
from keras.layers import Bidirectional
from keras.layers import CuDNNLSTM
from keras.layers import CuDNNGRU
from keras.layers import Layer

from keras import initializers
from keras import regularizers
from keras import constraints

from keras.preprocessing.text import Tokenizer
from keras.preprocessing.sequence import pad_sequences


class Attention(Layer):
    def __init__(self, step_dim,
                 W_regularizer=None, b_regularizer=None,
                 W_constraint=None, b_constraint=None,
                 bias=True, **kwargs):
        """
        Keras Layer that implements an Attention mechanism for temporal data.
        Supports Masking.
        Follows the work of Raffel et al. [https://arxiv.org/abs/1512.08756]
        # Input shape
            3D tensor with shape: `(samples, steps, features)`.
        # Output shape
            2D tensor with shape: `(samples, features)`.
        :param kwargs:
        Just put it on top of an RNN Layer (GRU/LSTM/SimpleRNN) with return_sequences=True.
        The dimensions are inferred based on the output shape of the RNN.
        Example:
            model.add(LSTM(64, return_sequences=True))
            model.add(Attention())
        """
        self.supports_masking = True
        #self.init = initializations.get('glorot_uniform')
        self.init = initializers.get('glorot_uniform')

        self.W_regularizer = regularizers.get(W_regularizer)
        self.b_regularizer = regularizers.get(b_regularizer)

        self.W_constraint = constraints.get(W_constraint)
        self.b_constraint = constraints.get(b_constraint)

        self.bias = bias
        self.step_dim = step_dim
        self.features_dim = 0
        super(Attention, self).__init__(**kwargs)

    def build(self, input_shape):
        assert len(input_shape) == 3

        self.W = self.add_weight((input_shape[-1],),
                                 initializer=self.init,
                                 name='{}_W'.format(self.name),
                                 regularizer=self.W_regularizer,
                                 constraint=self.W_constraint)
        self.features_dim = input_shape[-1]

        if self.bias:
            self.b = self.add_weight((input_shape[1],),
                                     initializer='zero',
                                     name='{}_b'.format(self.name),
                                     regularizer=self.b_regularizer,
                                     constraint=self.b_constraint)
        else:
            self.b = None

        self.built = True

    def compute_mask(self, input, input_mask=None):
        # do not pass the mask to the next layers
        return None

    def call(self, x, mask=None):
        # eij = K.dot(x, self.W) TF backend doesn't support it

        # features_dim = self.W.shape[0]
        # step_dim = x._keras_shape[1]

        features_dim = self.features_dim
        step_dim = self.step_dim

        eij = K.reshape(K.dot(K.reshape(x, (-1, features_dim)), K.reshape(self.W, (features_dim, 1))), (-1, step_dim))

        if self.bias:
            eij += self.b

        eij = K.tanh(eij)

        a = K.exp(eij)

        # apply mask after the exp. will be re-normalized next
        if mask is not None:
            # Cast the mask to floatX to avoid float64 upcasting in theano
            a *= K.cast(mask, K.floatx())

        # in some cases especially in the early stages of training the sum may be almost zero
        a /= K.cast(K.sum(a, axis=1, keepdims=True) + K.epsilon(), K.floatx())

        a = K.expand_dims(a)
        weighted_input = x * a
    #print weigthted_input.shape
        return K.sum(weighted_input, axis=1)

    def compute_output_shape(self, input_shape):
        #return input_shape[0], input_shape[-1]
        return input_shape[0],  self.features_dim


def threshold_search(y_true, y_predict):
    best_threshold = 0
    best_score = 0
    for t in [idx * 0.01 for idx in range(100)]:
        score = f1_score(y_true=y_true, y_pred=y_predict > t)
        if score > best_score:
            best_threshold = t
            best_score = score
    return float(best_threshold)


def build_vocab(sentences):
    """
    :param sentences: list of list of words
    :return: dictionary of words and their count
    """
    vocab = {}
    for sentence in sentences:
        for word in sentence:
            try:
                vocab[word] += 1
            except KeyError:
                vocab[word] = 1
    return vocab


def check_coverage(vocab, embeddings_idx):
    a = {}
    oov = {}
    k = 0
    idx = 0
    for word in vocab:
        try:
            a[word] = embeddings_idx[word]
            k += vocab[word]
        except:
            try:
                a[word] = embeddings_idx[word.capitalize()]
                k += vocab[word]
            except:
                oov[word] = vocab[word]
                idx += vocab[word]
                pass

    print('Found embeddings for {:.2%} of vocab'.format(len(a) / len(vocab)))
    print('Found embeddings for  {:.2%} of all text'.format(k / (k + idx)))
    sorted_x = sorted(oov.items(), key=operator.itemgetter(1))[::-1]

    return sorted_x


def clean_text(x):
    x = str(x)
    for punct in "/-'":
        x = x.replace(punct, ' ')
    for punct in '&':
        x = x.replace(punct, f' {punct} ')
    for punct in '?!.,"#$%\'()*+-/:;<=>@[\\]^_`{|}~' + '“”’':
        x = x.replace(punct, '')
    return x


def clean_numbers(x):
    x = re.sub('[0-9]{5,}', '#####', x)
    x = re.sub('[0-9]{4}', '####', x)
    x = re.sub('[0-9]{3}', '###', x)
    x = re.sub('[0-9]{2}', '##', x)
    return x


def replace_typical_misspell(text):
    def replace(match):
        return mispellings[match.group(0)]
    return mispellings_re.sub(replace, text)


def _get_mispell(mispell_dict):
    mispell_re = re.compile('(%s)' % '|'.join(mispell_dict.keys()))
    return mispell_dict, mispell_re


def build_classifier_rnn(layer_size, input_dim, dropout, embed_size, embedding_matrix, maxlen, units):
    inp = Input(shape=(maxlen,))
    
    emb_list = []
    for i in range(len(embedding_matrix)):
        emb = Embedding(input_dim, embed_size, weights=[embedding_matrix[i]], trainable=False)(inp)
        emb_list.append(emb)

    emb_concat_axis = 1
    if len(emb_list) > 1:
        x = Concatenate(axis=emb_concat_axis)(emb_list)
    else:
        x = emb_list[0]
    
    x = Bidirectional(CuDNNGRU(units, return_sequences=True, kernel_initializer='glorot_uniform'))(x)
    x = Bidirectional(CuDNNGRU(units, return_sequences=True, kernel_initializer='glorot_uniform'))(x)
    x = Bidirectional(CuDNNGRU(units, return_sequences=True, kernel_initializer='glorot_uniform'))(x)
    x = Attention(maxlen)(x)
    
    for index in range(0, len(layer_size)):
        x = Dense(units=layer_size[index], init='he_uniform', activation='relu')(x)
        x = BatchNormalization()(x)
        if dropout > 0:
            x = Dropout(dropout)(x)

    outp = Dense(units=1, init='glorot_uniform', activation='sigmoid')(x)

    model = Model(inputs=inp, outputs=outp)
    model.compile(loss='binary_crossentropy', optimizer='rmsprop', metrics=[f1])
    
    return model


def f1(y_true, y_pred):
    '''
    metric from here 
    https://stackoverflow.com/questions/43547402/how-to-calculate-f1-macro-in-keras
    '''
    def recall(y_true, y_pred):
        """Recall metric.

        Only computes a batch-wise average of recall.

        Computes the recall, a metric for multi-label classification of
        how many relevant items are selected.
        """
        true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))
        possible_positives = K.sum(K.round(K.clip(y_true, 0, 1)))
        recall = true_positives / (possible_positives + K.epsilon())
        return recall

    def precision(y_true, y_pred):
        """Precision metric.

        Only computes a batch-wise average of precision.

        Computes the precision, a metric for multi-label classification of
        how many selected items are relevant.
        """
        true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))
        predicted_positives = K.sum(K.round(K.clip(y_pred, 0, 1)))
        precision = true_positives / (predicted_positives + K.epsilon())
        return precision
    precision = precision(y_true, y_pred)
    recall = recall(y_true, y_pred)
    return 2*((precision*recall)/(precision+recall+K.epsilon()))


def clean_special_chars(text, punct, mapping):
    for p in mapping:
        text = text.replace(p, mapping[p])
    
    for p in punct:
        text = text.replace(p, f' {p} ')
    
    # Other special characters that I have to deal with in last    
    specials = {"вЂ™": "'",'\u200b': ' ', '…': ' ... ', '\ufeff': '', 'करना': '', 'है': ''}
    for s in specials:
        text = text.replace(s, specials[s])
    
    return text


def clear_stopwords(text, stopwords):
    text_arr = text.split()
    text_arr = [word for word in text_arr if not word in set(stopwords)]
    text = ' '.join(text_arr)
    return text

mispell_dict = {"ain't": "is not", "aren't": "are not", "can't": "cannot", "'cause": "because",
                "could've": "could have", "couldn't": "could not", "didn't": "did not", "doesn't": "does not",
                "don't": "do not", "hadn't": "had not", "hasn't": "has not", "haven't": "have not", "he'd": "he would",
                "he'll": "he will", "he's": "he is", "how'd": "how did", "how'd'y": "how do you", "how'll": "how will",
                "how's": "how is", "I'd": "I would", "I'd've": "I would have", "I'll": "I will",
                "I'll've": "I will have", "I'm": "I am", "I've": "I have", "i'd": "i would", "i'd've": "i would have",
                "i'll": "i will", "i'll've": "i will have", "i'm": "i am", "i've": "i have", "isn't": "is not",
                "it'd": "it would", "it'd've": "it would have", "it'll": "it will", "it'll've": "it will have",
                "it's": "it is", "let's": "let us", "ma'am": "madam", "mayn't": "may not", "might've": "might have",
                "mightn't": "might not", "mightn't've": "might not have", "must've": "must have", "mustn't": "must not",
                "mustn't've": "must not have", "needn't": "need not", "needn't've": "need not have",
                "o'clock": "of the clock", "oughtn't": "ought not", "oughtn't've": "ought not have",
                "shan't": "shall not", "sha'n't": "shall not", "shan't've": "shall not have", "she'd": "she would",
                "she'd've": "she would have", "she'll": "she will", "she'll've": "she will have", "she's": "she is",
                "should've": "should have", "shouldn't": "should not", "shouldn't've": "should not have",
                "so've": "so have", "so's": "so as", "this's": "this is", "that'd": "that would",
                "that'd've": "that would have", "that's": "that is", "there'd": "there would",
                "there'd've": "there would have", "there's": "there is", "here's": "here is", "they'd": "they would",
                "they'd've": "they would have", "they'll": "they will", "they'll've": "they will have",
                "they're": "they are", "they've": "they have", "to've": "to have", "wasn't": "was not",
                "we'd": "we would", "we'd've": "we would have", "we'll": "we will", "we'll've": "we will have",
                "we're": "we are", "we've": "we have", "weren't": "were not", "what'll": "what will",
                "what'll've": "what will have", "what're": "what are", "what's": "what is", "what've": "what have",
                "when's": "when is", "when've": "when have", "where'd": "where did", "where's": "where is",
                "where've": "where have", "who'll": "who will", "who'll've": "who will have", "who's": "who is",
                "who've": "who have", "why's": "why is", "why've": "why have", "will've": "will have",
                "won't": "will not", "won't've": "will not have", "would've": "would have", "wouldn't": "would not",
                "wouldn't've": "would not have", "y'all": "you all", "y'all'd": "you all would",
                "y'all'd've": "you all would have", "y'all're": "you all are", "y'all've": "you all have",
                "you'd": "you would", "you'd've": "you would have", "you'll": "you will", "you'll've": "you will have",
                "you're": "you are", "you've": "you have", 'colour': 'color', 'centre': 'center',
                'favourite': 'favorite', 'travelling': 'traveling', 'counselling': 'counseling', 'theatre': 'theater',
                'cancelled': 'canceled', 'labour': 'labor', 'organisation': 'organization', 'wwii': 'world war 2',
                'citicise': 'criticize', 'youtu ': 'youtube ', 'Qoura': 'Quora', 'sallary': 'salary', 'Whta': 'What',
                'narcisist': 'narcissist', 'howdo': 'how do','Howdo': 'how do', 'whatare': 'what are', 'howcan': 'how can',
                'howmuch': 'how much', 'howmany': 'how many', 'whydo': 'why do', 'doI': 'do I', 'theBest': 'the best',
                'howdoes': 'how does', 'mastrubation': 'masturbation', 'mastrubate': 'masturbate', 'masterbation': 'masturbation',
                "mastrubating": 'masturbating', 'pennis': 'penis', 'Etherium': 'Ethereum', 'narcissit': 'narcissist',
                'bigdata': 'big data', '2k17': '2017', '2k18': '2018', 'qouta': 'quota', 'exboyfriend': 'ex boyfriend',
                'airhostess': 'air hostess', "whst": 'what', 'watsapp': 'social medium', 'demonitisation': 'demonetization',
                'demonitization': 'demonetization', 'demonetisation': 'demonetization','didnt': 'did not','doesnt': 'does not',
                'isnt': 'is not','Shouldnt': 'should not', 'shouldnt':'should not','instagram': 'social medium','whatsapp': 'social medium',
                'snapchat': 'social medium','facebook': 'social medium', 'cos2': 'cosine 2', 'miui': 'Mi UI', 'Wasnt': 'was not',
                'Pornhub': 'Porn hub', 'optimisation': 'optimization', 'modernisation': 'modernization', 'Devops': 'Dev OPS',
                'cosx': 'cosine x', 'Doesnt': 'does not', 'Isnt': 'is not', 'Whatis': 'what is', 'wasnt': 'was not', 'recognise': 'recognize',
                'modelling': 'modeling', 'judgement': 'judgment', 'hasnt': 'has not', 'WeChat': 'social medium', 'analyse': 'analyze',
                'programmes': 'programs', 'specialisation': 'specialization', 'organised': 'organized', 'downvote': 'down vote',
                'sinx': 'sine x', 'aeroplane': 'plane', 'honour': 'honor', 'Couldn': 'could not', 'memorise': 'memorize', 'flavour': 'flavor',
                'masterbating': 'masturbating','masterbate': 'masturbate', 'downvoted': 'down voted', 'flavours': 'flavors', 'sin2x': 'sine 2 x', 'downvoting': 'down voting',
                'specialise': 'specialize', 'cryptocurrencies': 'crypto currencies', 'programme': 'program', 'Snapchat': 'social medium',
                'realise': 'realize', 'upvotes': 'up votes', 'upvoted': 'up voted', 'upvote': 'up vote', 'Paytm': 'pay time',
                'cryptocurrency': 'crypto currency', 'Cryptocurrency': 'crypto currency', 'bitcoins': 'crypto currency','bitcoin': 'crypto currency',
                'Instagram': 'social medium', 'Whatsapp': 'social medium', 'WhatsApp': 'social medium', 'Bitcoins': 'crypto currency', 
                'Bitcoin': 'crypto currency', 'Facebook': 'social medium', 'programrs': 'programmers', 'programr': 'programmer',
                'Programrs': 'programmers', 'Programr': 'programmer', 'civilisation': 'civilization','selfies': 'photo', 'selfie': 'photo',
                'globalisation': 'globalization', 'upvoting': 'up voting', 'litecoin': 'crypto currency', 'femdom': 'porn', 'crossdress': 'porn',
                'Howcan': 'how can', 'Howmany': 'how many', 'WeWork': 'we work', 'Cryptocurrencies': 'crypto currencies', 'Quorans': 'users',
                'Brexit': 'Britain exit', 'Coursera': 'online university', 'Blockchain': 'block chain', 'blockchain': 'block chain', 'Redmi': 'smartphone',
                'btech': 'bachelor of technology', 'Btech': 'bachelor of technology', 'Ethereum': 'block chain platform', 'ethereum': 'block chain platform',
                'Lyft': 'transportation service', 'vape': 'electronic cigarette', 'neurotypicals': 'neurotypical', 'altcoins': 'crypto currency', 'Litecoin': 'crypto currency',
                'Altcoins': 'crypto currency', 'litecoin': 'crypto currency', 'cisgender': 'matching gender', 'Whatare': 'what are', 'MeToo': 'me too', 'metoo': 'me too', 
                'friendzone' : 'friend zone', 'BTECH': 'bachelor of technology', 'psycopath': 'psychopath', 'Xiaomi': 'chinese electronics company',
                'Fortnite': 'video game', 'Fiverr': 'platform for freelancers', 'Pinterest': 'social media', 'hairfall': 'hair loss',
                'emojis': 'smile', 'altcoin': 'crypto currency', 'councelling': 'counseling', 'Truecaller': 'smartphone application',
                'BREXIT': 'Britain exit', 'athiesm': 'atheism', 'beleifs': 'beliefs', 'elecric' : 'electric', 'cleanshot' : 'clean shot',
                'rutine': 'routine', 'overDoes': 'over does', 'coinmarketcap': 'coin market cap', 'frwquency': 'frequency'
}

punct = "/-'?!.,#$%\'()*+-/:;<=>@[\\]^_`{|}~" + '""“”’' + '∞θ÷α•à−β∅³π‘₹´°£€\×™√²—–&'

punct_mapping = {
                "‘": "'", "₹": "e", "´": "'", "°": "", "€": "e", "™": "tm", "√": " sqrt ", "×": "x", "²": "2", "—": "-", 
                "–": "-", "’": "'", "_": "-", "`": "'", '“': '"', '”': '"', '“': '"', "£": "e", '∞': 'infinity', 
                'θ': 'theta', '÷': '/', 'α': 'alpha', '•': '.', 'à': 'a', '−': '-', 'β': 'beta', '∅': '', '³': '3', 'π': 'pi'
}

mispellings, mispellings_re = _get_mispell(mispell_dict)

    
def load_google_news(vocab, max_features):
    news_path = '../input/embeddings/GoogleNews-vectors-negative300/GoogleNews-vectors-negative300.bin'
    embeddings_index = KeyedVectors.load_word2vec_format(news_path, binary=True)
    words_not_in_embeddings = []
    nb_words = min(max_features, len(vocab))
    embedding_matrix = np.zeros((nb_words, embed_size))
    for word, i in vocab.items():
        if i >= max_features:
            continue
        try:
            embedding_vector = embeddings_index[word]
            if embedding_vector is not None:
                embedding_matrix[i] = embedding_vector
        except:
            try:
                embedding_vector = embeddings_index[word.capitalize()]
                if embedding_vector is not None:
                    embedding_matrix[i] = embedding_vector
            except:
                words_not_in_embeddings.append(word)
    return embedding_matrix


train = pd.read_csv("../input/train.csv")
test = pd.read_csv("../input/test.csv")
print("Train shape : ", train.shape)
print("Test shape : ", test.shape)

stopwords = ['a','to','of','and', 'A', 'To', 'And', 'Of', '...', '-',"'",'/']

train["question_text"] = train["question_text"].apply(lambda x: clean_text(x))
train["question_text"] = train["question_text"].apply(lambda x: clean_numbers(x))
train["question_text"] = train["question_text"].apply(lambda x: clean_special_chars(x, punct, punct_mapping))
train["question_text"] = train["question_text"].apply(lambda x: replace_typical_misspell(x))
train["question_text"] = train["question_text"].apply(lambda x: replace_typical_misspell(x))
train["question_text"] = train["question_text"].apply(lambda x: clear_stopwords(x, stopwords))

test["question_text"] = test["question_text"].apply(lambda x: clean_text(x))
test["question_text"] = test["question_text"].apply(lambda x: clean_numbers(x))
test["question_text"] = test["question_text"].apply(lambda x: clean_special_chars(x, punct, punct_mapping))
test["question_text"] = test["question_text"].apply(lambda x: replace_typical_misspell(x))
test["question_text"] = test["question_text"].apply(lambda x: replace_typical_misspell(x))
test["question_text"] = test["question_text"].apply(lambda x: clear_stopwords(x, stopwords))

train_X = train["question_text"].values
test_X = test["question_text"].values
train_y = train['target'].values

x_train, x_valid, y_train, y_valid = train_test_split(train_X, train_y, test_size=0.02, random_state=42)

embed_size = 300
max_features = 90000
maxlen = 50

tokenizer = Tokenizer(num_words=max_features)
tokenizer.fit_on_texts(list(np.concatenate((train_X, test_X), axis=0)))

x_train = tokenizer.texts_to_sequences(x_train)
x_valid = tokenizer.texts_to_sequences(x_valid)
test_X = tokenizer.texts_to_sequences(test_X)

x_train = pad_sequences(x_train, maxlen=maxlen)
x_valid = pad_sequences(x_valid, maxlen=maxlen)
test_X = pad_sequences(test_X, maxlen=maxlen)

embedding_matrix_0 = load_google_news(tokenizer.word_index, max_features)
print('Google news embedding loaded')

loss_threshold = 0.1070
model_count = 2
models = []
for i in range(0, model_count):
    print('Training model ' + str(i + 1) + '/' + str(model_count))
    
    clf = build_classifier_rnn(
        layer_size=[128,32],
        input_dim=max_features,
        dropout=0.2,
        embed_size=embed_size,
        embedding_matrix=[embedding_matrix_0],
        maxlen=maxlen,
        units = 64)
    
    history = clf.fit(
        x_train,
        y_train,
        callbacks=[
            EarlyStopping(
                monitor='val_loss',
                mode='min',
                restore_best_weights=True,
                patience=2)
        ],
        batch_size=512,
        epochs=30,
        validation_data=(x_valid, y_valid),
        verbose=2,
        shuffle=True)
    
    true_model = False
    for val_loss in history.history['val_loss']:
        if val_loss < loss_threshold:
            true_model = True
            break
    
    if true_model:
        models.append(clf)

print('Complete fitting. Selected ' + str(len(models)) + ' models. Loss threshold ' + str(loss_threshold) + '.')

models[0].summary()

predictions = np.zeros((len(x_valid), len(models)))
for m in range(0, len(models)):
    pred = models[m].predict(x_valid)
    for index in range(0, len(pred)):
        predictions[index][m] = pred[index]

stacker= LinearRegression()
stacker.fit(predictions, y_valid)

coefs = stacker.coef_

print('Ensemble coefficients:')
print(coefs)

y_pred = stacker.predict(predictions)

threshold = threshold_search(y_valid, y_pred)

print('Using threshold: ' + str(threshold))

accuracy = accuracy_score(y_valid, y_pred > threshold)

print('Tranning accuracy: ' + str(round(accuracy, 3)))

f_score = f1_score(y_valid, y_pred > threshold)

print('Tranning F1 score: ' + str(round(f_score, 3)))

print('+------------------------------------+')
print('| SAVE SUBMISSION                    |')
print('+------------------------------------+')

predictions = np.zeros((len(test_X), len(models)))
for m in range(0, len(models)):
    pred = models[m].predict(test_X)
    for index in range(0, len(pred)):
        predictions[index][m] = pred[index]

y_pred = stacker.predict(predictions)

out_df = pd.DataFrame(columns=['qid', 'prediction'])
out_df['qid'] = test["qid"].values
out_df['prediction'] = (y_pred > threshold).astype(int)
out_df.to_csv("submission.csv", index=False)

print('submission.csv file created')
