# Data preparation --------------------------------------------------------

set.seed(1234)

Pooled_db <- readRDS(file = file.path(cleaned_data_dir, "internet_use_prediction", "SEDLAC_WAEMU_PIAAC.rds"))

model_db <- Pooled_db %>% 
  filter(!is.na(internet_use)) %>% 
  mutate(internet_use = factor(internet_use)) %>% 
  select(-countrycode, -countryname, -year, -weight, -weight_income, -has_a_job, -labor_status, -wage_hourly, -wage_monthly, -welfare, -welfare_decile, -age, -urban, -age_group_16_25, -edu_lvl_0, -ISIC_raw, -ISCO_1d_0, -starts_with("ISIC_1d_"), -industrycat10_1,  -literacy_rate, -`region_North America`, -`region_South Asia`, -starts_with("region"), -starts_with("income"), -ISCO_1d, -ISCO_2d, -ISCO_3d, -ISCO_4d) %>% 
  drop_na()

# Sample selection

model_db <- model_db %>% 
  filter(!age_group_66_more == 1)

db_split <- initial_split(model_db, prop = 0.7)

train <- training(db_split) %>% 
  mutate(weight = importance_weights(weight_rescaled))

test <- testing(db_split)

# Folds  ------------------------------------------------------------------

set.seed(1234)

kfolds <- vfold_cv(train, v = 5)

# Recipes for model estimation --------------------------------------------

recipe_interact <- 
  recipe(internet_use ~ ., data = train) %>%
  step_rm(weight_rescaled) %>% 
  step_interact(terms = ~ gender:starts_with("age_group")) %>% # gender x age group
  step_interact(terms = ~ gender:starts_with("edu_lvl")) %>% # gender x education level
  step_interact(terms = ~ gender:starts_with("ISCO")) %>% # gender x occupation
  step_interact(terms = ~ starts_with("age_group"):starts_with("edu_lvl")) %>% # age group x education level
  step_zv(all_predictors()) %>% 
  step_lincomb(all_predictors()) %>%
  step_corr(all_predictors())

# Defining models

logit_lasso <- logistic_reg(penalty = tune(), mixture = 1) %>%
  set_engine("glmnet")

# Logit-lasso model -------------------------------------------------------

workflow_logit_lasso <- workflow() %>% 
  add_recipe(recipe_interact) %>% 
  add_model(logit_lasso) %>% 
  add_case_weights(weight)

# Tune grid

lambda_grid_lasso <- grid_regular(penalty(), levels = 50)

registerDoParallel()

tune_lasso <- workflow_logit_lasso %>% 
  tune_grid(resamples = kfolds,
            grid = lambda_grid_lasso,
            metrics = metric_set(f_meas),
            control=control_grid(verbose=TRUE))

results_tuning_lasso <- tune_lasso %>%
  collect_metrics()

stopImplicitCluster()

# Best fit

tune_best_lasso <- tune_lasso %>% 
  select_best(metric = "f_meas")

best_lasso_model <- logistic_reg(penalty = tune_best_lasso$penalty, mixture = 1) %>%
  set_engine("glmnet")

workflow_logit_lasso <- workflow() %>% 
  add_recipe(recipe_interact) %>% 
  add_model(best_lasso_model) %>% 
  add_case_weights(weight)

best_lasso_fit <- workflow_logit_lasso %>% 
  fit(train)

# Predictions

logit_lasso_predictions <- as.data.frame(list(internet_prob_predict = predict(best_lasso_fit, new_data = test, type = "prob")$.pred_1, internet_use = test$internet_use)) 

roc_thresh_logit_lasso <- coords(roc(logit_lasso_predictions, internet_use, internet_prob_predict), x = "best", best.method = "closest.topleft") # optimal threshold = 0.693

logit_lasso_predictions <- logit_lasso_predictions %>% 
  mutate(internet_predict = ifelse(internet_prob_predict > roc_thresh_logit_lasso$threshold, 1, 0),
         internet_predict = factor(internet_predict, levels = c(1,0)),
         internet_use = factor(internet_use, levels = c(1,0)))

eval_metrics_model_logit_lasso <- bind_rows(
  metric_set(accuracy, precision, recall, f_meas)(logit_lasso_predictions, truth = internet_use, estimate = internet_predict), 
  roc_auc(logit_lasso_predictions, internet_use, internet_prob_predict))

# Exporting predictions

Logit_lasso_predictions_db <- Pooled_db %>% 
  mutate(internet_prob_predict = predict(best_lasso_fit, new_data = Pooled_db, type = "prob")$.pred_1,
         internet_use_predict  = ifelse(internet_prob_predict > roc_thresh_logit_lasso$threshold, 1, 0))

saveRDS(Logit_lasso_predictions_db, file = file.path(cleaned_data_dir, "internet_use_prediction", "Pooled_logit_lasso_predictions.rds"))

# Prediction on GLD -------------------------------------------------------

GLD_db <- readRDS(file = file.path(cleaned_data_dir, "internet_use_prediction", "GLD_EAPCE_others.rds"))

# Predictions from logit model (generating predictions only for the best model)

GLD_db <- GLD_db %>%
  mutate(internet_prob_predict = predict(best_lasso_fit, new_data = GLD_db, type = "prob")$.pred_1,
         internet_use_predict  = ifelse(internet_prob_predict > roc_thresh_logit_lasso$threshold, 1, 0))

saveRDS(GLD_db, file = file.path(cleaned_data_dir, "internet_use_prediction", "GLD_logit_lasso_predictions.rds"))
