import string
import sys
import warnings

import lightgbm as lgb
from scipy.sparse import hstack, vstack
from sklearn.feature_extraction.text import TfidfVectorizer

from cetune.tune_util import *
 
warnings.filterwarnings('ignore')
start_time = int(time.time())
end_time = start_time + 9 * 3600 - 60 * 60

num_round = 1000


def get_values(x):
    return x.values if hasattr(x, 'values') else x


class LgbTrainer:
    def __init__(self, params):
        self.params = params
        self.model = None

    def set_params(self, **params):
        self.params.update(params)

    def fit(self, x, y):
        self.model = lgb.train(self.params, lgb.Dataset(get_values(x), label=get_values(y)),
                               num_boost_round=self.params['num_boost_round'])
        return self

    def predict(self, x):
        return self.model.predict(get_values(x))


def f1(y, p, detail=True):
    f1_pairs = sorted([(p_th, metrics.f1_score(y, p > p_th)) for p_th in np.arange(0.1, 0.91, 0.01)],
                      key=lambda pair: (pair[1], pair[0]), reverse=True)[:5]
    f1s = [f1_score for p_th, f1_score in f1_pairs]
    f1_mean = np.mean(f1s)
    if detail:
        print(f'f1_mean={f1_mean}, f1_std={np.std(f1s)}, {f1_pairs}')
    return f1_mean


@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 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()

    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 tune_params(data_dir, tune_dir='.'):
    x, y, ts_x, submission = get_data(data_dir)
    del ts_x, submission
    gc.collect()

    params = {'objective': 'binary', 'metric': 'auc', 'verbose': -1, 'nthread': 4, 'num_boost_round': num_round,
              'learning_rate': 0.1, 'num_leaves': 31, 'scale_pos_weight': 10.0, 'max_depth': 0, 'min_data': 20,
              'bagging_fraction': 0.8, 'feature_fraction': 0.8, 'bagging_freq': 1, 'lambda_l1': 0, 'lambda_l2': 0}
    model = LgbTrainer(params=params)

    init_param = [('learning_rate', 0.1), ('num_leaves', 31), ('scale_pos_weight', 10.0), ('max_depth', 0),
                  ('min_data', 20), ('bagging_fraction', 0.8), ('feature_fraction', 0.8), ('lambda_l1', 0),
                  ('lambda_l2', 0)]
    param_dic = {'learning_rate': [.04, .08, .02],
                 'num_leaves': [20, 40, 80],
                 'max_depth': [0, 8, 10, 12, 14],
                 'scale_pos_weight': [1.0, 2, 4, 8, 16],
                 'min_data': [20, 40, 80, 160],
                 'bagging_fraction': ['grid', .1, .2, .3, .4, .5, .6, .7, .8, .9, 1.0],
                 'feature_fraction': ['grid', .1, .2, .3, .4, .5, .6, .7, .8, .9, 1.0],
                 'lambda_l1': [0.0, .01, .02, .04, .08, .2, .4, .8],
                 'lambda_l2': [0.0, .01, .02, .04, .08, .2, .4, .8]}

    tune(model, (x, y), init_param, param_dic, measure_func=metrics.roc_auc_score, detail=True, random_state=1000, data_dir=tune_dir,
         kc=(3, 1), score_min_gain=2e-4, kfold_func=kfold, task_id=f'lgb_{num_round}',
         non_ordinal_factors=['max_depth'], end_time=end_time)
 
 
def run():
    sys.stderr = sys.stdout = open(os.path.join(f'log.txt'), 'w')

    print(os.listdir('../input'))
    data_dir = '../input/quora-insincere-questions-classification'
    print(os.listdir(data_dir))
 
    task_id = f'lgb_{num_round}'
    cache_root_dir = f'../input/quora-lgb-tune-param-reboot'
    print(os.listdir(cache_root_dir))
    shutil.copytree(os.path.join(cache_root_dir, 'cache', task_id), os.path.join('.', 'cache', task_id))

    tune_params(data_dir)


if __name__ == '__main__':
    run()
