library(data.table)
library(Matrix)
library(xgboost)
library(caret)
library(dplyr)

cat("Init")
set.seed(12345)
PATH <- "../input/"

cat("Load data")
train <- fread(paste0(PATH,"train.csv"), sep=",", na.strings = "", stringsAsFactors=T)
transactions <- fread(paste0(PATH,"transactions.csv"), sep=",", na.strings = "", stringsAsFactors=T)
members <- fread(paste0(PATH,"members.csv"), sep=",", na.strings = "", stringsAsFactors=T)
sample_submission_zero <- fread(paste0(PATH,"sample_submission_zero.csv"), sep=",", na.strings = "", stringsAsFactors=T)

cat("Combine train and test files")
sample_submission_zero$is_churn <- NA
data <- rbind(train, sample_submission_zero)
data[,is_duplicate := as.numeric(duplicated(as.character(data$msno)) | duplicated(as.character(data$msno),fromLast=T))]
rm(train);gc()

cat("Format gender and remove NA's")
members[,gender := as.numeric(gender)]
members$gender[is.na(members$gender)] <- 0

cat("Format dates and do some feature engineering")
members[,":="(reg_fulldate = members$registration_init_time
             ,registration_init_time = as.Date(as.character(registration_init_time), '%Y%m%d')
             ,exp_fulldate = expiration_date
             ,expiration_date = as.Date(as.character(expiration_date), '%Y%m%d'))]
members[,":="(reg_year = year(registration_init_time)
             ,reg_month = month(registration_init_time)
             ,reg_mday = mday(registration_init_time)
             ,reg_wday = wday(registration_init_time)
             ,exp_year = year(expiration_date)
             ,exp_month = month(expiration_date)
             ,exp_mday = mday(expiration_date)
             ,exp_wday = wday(expiration_date)
             ,date_diff = as.numeric(expiration_date - registration_init_time))]
members <- subset(members, select = -c(registration_init_time, expiration_date))

cat("Merge data and members")
data <- merge(data, members, by = "msno", all.x = TRUE)
rm(members);gc()

cat("Reduce size of transactions a bit")
transactions <- transactions[transactions$msno %in% levels(data$msno),]

cat("Get amount of transactions per user")
transactions[,n_transactions := .N, by = msno]

cat("Get difference between plan price and payment amount")
transactions[,payment_price_diff := plan_list_price - actual_amount_paid]

cat("Aggregate by user, get mean of columns")
cat("I don't think the transaction dates are useful for now, so let's remove them")
transactions <- transactions[,lapply(.SD,mean,na.rm=T), by = msno, .SDcols = names(transactions)[c(2:6,9:11)]]

cat("Merge data and transactions")
data <- merge(data, transactions, by = "msno", all.x = TRUE)
rm(transactions);gc()

cat("Prepare for xgb")
cvFolds <- createFolds(data$is_churn[!is.na(data$is_churn)], k=5, list=TRUE, returnTrain=FALSE)
varnames <- setdiff(colnames(data), c("msno", "is_churn"))
train_sparse <- Matrix(as.matrix(data[!is.na(is_churn), varnames, with=F]), sparse=TRUE)
test_sparse <- Matrix(as.matrix(data[is.na(is_churn), varnames, with=F]), sparse=TRUE)
y_train <- data[!is.na(is_churn),is_churn]
test_ids <- data[is.na(is_churn),msno]
dtrain <- xgb.DMatrix(data=train_sparse, label=y_train)
dtest <- xgb.DMatrix(data=test_sparse)

cat("Params for xgb")
param <- list(booster="gbtree",
              objective="binary:logistic",
              eval_metric="logloss",
              eta = .02,
              gamma = 1,
              max_depth = 6,
              min_child_weight = 1,
              subsample = .8,
              colsample_bytree = .8
)

cat("Xgb cross-validation, uncomment when running locally")
# xgb_cv <- xgb.cv(data = dtrain,
#                  params = param,
#                  nrounds = 1000,
#                  maximize = FALSE,
#                  prediction = TRUE,
#                  folds = cvFolds,
#                  print_every_n = 10,
#                  early_stopping_round = 50)
# best_iter <- xgb_cv$best_iteration
best_iter <- 1512

cat("xgb model")
xgb_model <- xgb.train(data = dtrain,
                       params = param,
                       watchlist = list(train = dtrain),
                       nrounds = best_iter,
                       verbose = 1,
                       print_every_n = 100
)

cat("Feature importance")
names <- dimnames(train_sparse)[[2]]
importance_matrix <- xgb.importance(names, model=xgb_model)
xgb.plot.importance(importance_matrix)

cat("Predict and output csv")
preds <- data.table(msno=test_ids, is_churn=predict(xgb_model,dtest))
preds <- merge(sample_submission_zero[,1], preds, by="msno", all.x=T, sort=F)
write.table(preds, "submission.csv", sep=",", dec=".", quote=FALSE, row.names=FALSE)

