import gc
import os
import string
import sys
import time
from contextlib import contextmanager
from datetime import datetime

import lightgbm as lgb
import numpy as np
import pandas as pd
from scipy.sparse import hstack, vstack
from sklearn import metrics
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.model_selection import ShuffleSplit

last_best_p_th = 0.5


@contextmanager
def timer(name):
    print(f'【{name}】 begin at 【{datetime.now().strftime("%Y-%m-%d %H:%M:%S")}】')
    t0 = time.time()
    yield
    print(f'【{name}】 done in 【{time.time() - t0:.0f}】 s')


def get_time_stamp():
    return datetime.now().strftime("%Y%m%d%H%M%S")


def f1(y, p, detail=False):
    global last_best_p_th

    left = 0.1
    right = 0.9
    pace = 0.01
    k = 5
    tol_times = 2 * k

    f1_pairs = []
    best_score = 0

    p_th = last_best_p_th
    sink_times = 0
    while p_th <= right:
        f1_score = metrics.f1_score(y, p > p_th)
        f1_pairs.append((round(p_th, 10), f1_score))
        if f1_score < best_score:
            sink_times += 1
        else:
            best_score = f1_score
            sink_times = 0
        if sink_times > tol_times:
            break
        p_th += pace

    p_th = last_best_p_th - pace
    sink_times = 0
    while p_th >= left:
        f1_score = metrics.f1_score(y, p > p_th)
        f1_pairs.append((round(p_th, 10), f1_score))
        if f1_score < best_score:
            sink_times += 1
        else:
            best_score = f1_score
            sink_times = 0
        if sink_times > tol_times:
            break
        p_th -= pace

    f1_pairs = sorted(f1_pairs, key=lambda pair: (pair[1], pair[0]), reverse=True)[:k]
    last_best_p_th = f1_pairs[0][0]
    f1s = [f1_score for p_th, f1_score in f1_pairs]
    f1_mean = np.mean(f1s)
    if detail:
        print(f'last_best_p_th={last_best_p_th}, f1_mean={f1_mean}, f1_std={np.std(f1s)}, {f1_pairs}')
    return f1_mean


def lgb_f1(p, train_data):
    return 'f1', f1(train_data.get_label(), p), True


def combine_features(features, batch_num=5):
    cols = []
    batch_size = features[0].shape[0] // batch_num + 1
    for i in range(batch_num):
        fts = [ft[i * batch_size: (i + 1) * batch_size] for ft in features]
        cols.append(hstack(fts, dtype=np.float32).tocsr())
    return vstack(cols)


def get_data(data_dir='../input'):
    with timer('load data'):
        train_df = pd.read_csv(os.path.join(data_dir, 'train.csv'))
        test_df = pd.read_csv(os.path.join(data_dir, 'test.csv'), index_col='qid')
        submission = pd.read_csv(os.path.join(data_dir, 'sample_submission.csv'))
        print(f'train_df: {train_df.shape}, test_df: {test_df.shape}, submission: {submission.shape}')
        test_df = submission.join(test_df, on='qid', how='inner').drop('prediction', axis=1)
        print(f'test_df: {test_df.shape}')
        gc.collect()

        train_df = train_df.fillna('the')
        test_df = test_df.fillna('the')
        gc.collect()

    with timer('encode text'):
        def encode_text(col):
            def count_chars(txt):
                _len = 0
                digit_cnt, number_cnt = 0, 0
                lower_cnt, upper_cnt, letter_cnt, word_cnt = 0, 0, 0, 0
                char_cnt, term_cnt = 0, 0
                conj_cnt, blank_cnt, punc_cnt = 0, 0, 0
                sign_cnt, marks_cnt = 0, 0

                flag = 10
                for ch in txt:
                    _len += 1
                    if ch in string.ascii_lowercase:
                        lower_cnt += 1
                        letter_cnt += 1
                        char_cnt += 1
                        if flag:
                            word_cnt += 1
                            if flag > 2:
                                term_cnt += 1
                            flag = 0
                    elif ch in string.ascii_uppercase:
                        upper_cnt += 1
                        letter_cnt += 1
                        char_cnt += 1
                        if flag:
                            word_cnt += 1
                            if flag > 2:
                                term_cnt += 1
                            flag = 0
                    elif ch in string.digits:
                        digit_cnt += 1
                        char_cnt += 1
                        if 1 != flag:
                            number_cnt += 1
                            if flag > 2:
                                term_cnt += 1
                            flag = 1
                    elif '_' == ch:
                        conj_cnt += 1
                        char_cnt += 1
                        if flag > 2:
                            term_cnt += 1
                        flag = 2
                    elif ch in string.whitespace:
                        blank_cnt += 1
                        flag = 3
                    elif ch in string.punctuation:
                        punc_cnt += 1
                        flag = 4
                    else:
                        sign_cnt += 1
                        if flag != 5:
                            marks_cnt += 1
                            flag = 5

                return (_len, digit_cnt, number_cnt, digit_cnt / (1 + number_cnt), lower_cnt, upper_cnt, letter_cnt,
                        word_cnt, letter_cnt / (1 + word_cnt), char_cnt, term_cnt, char_cnt / (1 + term_cnt), conj_cnt,
                        blank_cnt, punc_cnt, sign_cnt, marks_cnt, sign_cnt / (1 + marks_cnt))

            return np.array(list(col.apply(count_chars)), dtype=np.uint16)

        tr_cnts = encode_text(train_df.question_text)
        ts_cnts = encode_text(test_df.question_text)
        gc.collect()
        print(f'tr_cnts: {tr_cnts.shape}, ts_cnts: {ts_cnts.shape}')

    with timer('vectorize text'):
        tvr = TfidfVectorizer(token_pattern=r'(?u)\w+|[^\w\s]', strip_accents='unicode', min_df=3)
        tr_x = tvr.fit_transform(train_df.question_text)
        ts_x = tvr.transform(test_df.question_text)
        print(f'tr_x: {tr_x.shape}, ts_x: {ts_x.shape}')

    with timer('combine features'):
        tr_x = combine_features([tr_x, tr_cnts])
        gc.collect()
        ts_x = combine_features([ts_x, ts_cnts])
        gc.collect()
        print(f'tr_x: {tr_x.shape}, ts_x: {ts_x.shape}')

    y = train_df.target.values.copy()
    del train_df, test_df
    gc.collect()

    return tr_x, y, ts_x, submission


def run_lgb(tr_x, y, ts_x, oof_seed=1):
    with timer('tv split'):
        tind, vind = next(ShuffleSplit(n_splits=1, test_size=0.15, random_state=oof_seed * 10000).split(y))
        tx = tr_x[tind]
        ty = y[tind]
        vx = tr_x[vind]
        vy = y[vind]
        print(f'tx: {tx.shape}, vx: {vx.shape}')
        print(f'ty(>0): {np.sum(ty)}, vy(>0): {np.sum(vy)}')
        del tr_x, y
        gc.collect()

    # params = {'objective': 'binary', 'metric': 'None', 'verbose': -1, 'nthread': 4,
    #           'learning_rate': 0.1, 'num_leaves': 31, 'max_depth': 0, 'min_data': 20, 'bagging_fraction': 0.8,
    #           'feature_fraction': 0.8, 'bagging_freq': 1, 'lambda_l1': 0, 'lambda_l2': 0, 'is_unbalance': True}
    params = {'objective': 'binary', 'metric': 'auc', 'verbose': -1, 'nthread': 4, 'scale_pos_weight': 2.0,
              'learning_rate': 0.1, 'num_leaves': 31, 'max_depth': 0, 'min_data': 20, 'bagging_fraction': 0.8,
              'feature_fraction': 0.8, 'bagging_freq': 1, 'lambda_l1': 0, 'lambda_l2': 0}
    # params.update([('learning_rate', 0.0875), ('num_leaves', 74), ('scale_pos_weight', 1.0), ('max_depth', 0),
    #                ('min_data', 15), ('bagging_fraction', 0.8), ('feature_fraction', 1.0), ('lambda_l1', 1.1),
    #                ('lambda_l2', 0.7)])
    # params.update([('learning_rate', 0.08125), ('num_leaves', 78), ('scale_pos_weight', 1.25), ('max_depth', 0),
    #               ('min_data', 6), ('bagging_fraction', 0.8), ('feature_fraction', 0.6), ('lambda_l1', 1.2),
    #               ('lambda_l2', 0.7)])
    # params.update([('learning_rate', 0.11), ('num_leaves', 91), ('scale_pos_weight', 1.75), ('max_depth', 0), 
    #                 ('min_data', 6), ('bagging_fraction', 0.9), ('feature_fraction', 0.6), ('lambda_l1', 1.7), 
    #                 ('lambda_l2', 0.08)])
                    
    # params.update([('learning_rate', 0.12), ('num_leaves', 40), ('scale_pos_weight', 2), ('max_depth', 0), 
    #                 ('min_data', 20), ('bagging_fraction', 0.8), ('feature_fraction', 0.5), ('lambda_l1', 1.7), 
    #                 ('lambda_l2', 0.02)])
    params.update([('learning_rate', 0.12), ('num_leaves', 40), ('scale_pos_weight', 2), ('max_depth', 0), 
                    ('min_data', 15), ('bagging_fraction', 0.8), ('feature_fraction', 0.9), ('lambda_l1', 1.1), 
                    ('lambda_l2', 0.9)])
 
    with timer('train'):
        model = lgb.train(params, lgb.Dataset(tx, label=ty), 100000, [lgb.Dataset(vx, label=vy)],
                          early_stopping_rounds=200, verbose_eval=200)#, feval=lgb_f1)
        model.save_model('lgb-1')
        print(f'best iteration: {model.best_iteration}')

    with timer('validation'):
        tp = model.predict(tx)
        t_auc = metrics.roc_auc_score(ty, tp)
        t_f1 = f1(ty, tp, True)

        vp = model.predict(vx)
        v_auc = metrics.roc_auc_score(vy, vp)
        v_f1 = f1(vy, vp, True)

        print(f't_auc: {t_auc}, v_auc: {v_auc}')
        print(f't_f1: {t_f1}, v_f1: {v_f1}')

    with timer('predict'):
        lgb_p = model.predict(ts_x)

    return lgb_p, last_best_p_th


def run():
    sys.stderr = sys.stdout = open(os.path.join('log.txt'), 'w')

    tr_x, y, ts_x, submission = get_data()
    lgb_p, lgb_p_th = run_lgb(tr_x, y, ts_x)

    submission['prediction'] = (lgb_p > lgb_p_th).astype(np.uint8)
    submission.to_csv('submission.csv', index=False)


if __name__ == '__main__':
    run()
