### Plot results from all models/mixtures
### Clear workspace
rm(list=ls(all=TRUE))
### Load Libraries
if (!require("pacman")) install.packages("pacman")
pacman::p_load(tidyverse, dplyr, reshape2, scoringRules, ggplot2, viridis, ggsci)
### Set working directory
setwd("D:/Toolkit/Backup/UGW-052141/C/Users/pedro/Documents/Projects/Fukushima/Code/Summary") #
### create some dplyr overrides
select = dplyr::select

# Get known mixture proportions
df.am <- read.csv("D:/Toolkit/Backup/UGW-052141/C/Users/pedro/Documents/Projects/Fukushima/Data/am_contributions.csv") %>%
  separate(sediment_sample, c("a", "Mix"), sep = 3) %>%
  select(-a) %>%
  mutate(Mix = as.numeric(Mix)) %>%
  melt(id.vars = "Mix") %>%
  tbl_df()  %>%
  rename(Source = variable)

# Read MixSIAR results
MS <- rbind(
  MS.lab <- read.csv("D:/Toolkit/Backup/UGW-052141/C/Users/pedro/Documents/Projects/Fukushima/MS/MS_Lab_2021.08.02--06.38.01/results.csv") %>%
    tbl_df() %>%
    mutate(Type = "LabMix"),
  MS.math <- read.csv("D:/Toolkit/Backup/UGW-052141/C/Users/pedro/Documents/Projects/Fukushima/MS/MS_Math_2021.08.02--11.46.07/results.csv") %>%
    tbl_df() %>%
    mutate(Type = "MathMix")
) %>%
  mutate(Model = "MS") %>%
  select(-variable)

# Read MVN results
MVN <- read.csv("D:/Toolkit/Backup/UGW-052141/C/Users/pedro/Documents/Projects/Fukushima/MVN_Mixing_Model_2021.07.30--10.23.09/ALL_MODEL_RESULTS_FROM_INDIVIDUAL_SAMPLES.csv") %>%
  select(-c(X)) %>%
  melt()  %>%
  separate(NA., c("a", "Mix"), sep = 3) %>%
  separate(variable, c("A", "Source")) %>%
  separate(Mix, c("Mix", "Type")) %>%
  mutate(B = paste0(A, "_",Source)) %>% 
  select(Mix, Type, value, B) %>%
  rename(Source = B) %>%
  tbl_df() %>%
  mutate(Type = ifelse(Type == "LM", "LabMix", "MathMix"),
         Model = "BMM")

# Combine models and mixtures
df.comb <- rbind(MS, MVN) %>%
  mutate(Source = ifelse(Source == "source_D", "Decontaminated",
                         ifelse(Source == "source_F", "Forest",
                                ifelse(Source == "source_PF", "Cropland", 
                                       "Subsurface")))) %>%
  group_by(Mix, Type, Source, Model) %>%
  summarise(q5 = quantile(value, 0.05),
            q25 = quantile(value, 0.25),
            median = median(value),
            q75 = quantile(value, 0.75),
            q95 = quantile(value, 0.95)) %>%
  melt(id.vars = c("Mix", "Type", "Source", "Model")) %>%
  tbl_df() %>%
  rename(Quant = variable) 

# Prepare data for plotting the lines/ribbons ribbons 
df.ribbon <- df.comb  %>%
  dcast(Source + Mix + Type + Model ~ Quant, value.var = "value") %>%
  tbl_df() %>%
  mutate(Mix = as.numeric(Mix))

df.obs <- df.am %>%
  rename(Observed = value) %>%
  mutate(Source = ifelse(Source == "source_D", "Decontaminated",
                         ifelse(Source == "source_F", "Forest",
                                ifelse(Source == "source_PF", "Cropland", 
                                       "Subsurface")))) %>%
  tbl_df() 

df.ribbon.comb <- df.ribbon %>% 
  inner_join(df.obs, by = c("Mix", "Source")) %>%
  tbl_df() %>%
  mutate(Error = (Observed - median)) %>%
  mutate(Mix = as.factor(Mix)) %>%
  mutate(Model = ifelse(Model == "MS", "MixSIAR", Model),
         Type = ifelse(Type == "LabMix", "Laboratory", "Mathematical")) 

# Figure 
ggplot(df.ribbon.comb %>%
         # filter(Type == "LabMix") %>%
         mutate(Mix = as.numeric(Mix)), aes(x = Mix)) +
  geom_line(aes(y = median, colour = Model, group = Model), size = 0.5) +
  geom_ribbon(aes(ymin = q25, ymax = q75, fill = Model, group = Model),
              alpha = 0.2) +
  scale_color_viridis(discrete = T, begin = .8, end = 0.2, option = "D") +
  scale_fill_viridis(discrete = T, begin = .8, end = 0.2, option = "D") +
  # scale_colour_manual(values = c("firebrick", "turquoise3")) +
  # scale_fill_manual(values = c("firebrick", "turquoise3")) +
  geom_line(aes(y = q5, colour = Model, group = Model), linetype = "dashed") +
  geom_line(aes(y = q95, colour = Model, group = Model), linetype = "dashed") +
  geom_point(aes(y = Observed), shape = 1) +
  facet_grid(rows = vars(Source), cols = vars(Type)) +
  theme_bw() +
  ylab("Source proportion\n") +
  xlab("\nMixture number") +
  theme(panel.spacing = unit(1, "lines")) +
  theme(axis.text.x = element_text(angle = 45, hjust = 1),
        text = element_text(size = 12),
        legend.position = "top",
        legend.title = element_blank()) 

#Save figure
ggsave("Models_Ribbons_Grid_Final.png", width = 10, height = 6.4)

require(hydroGOF)

# Calculate summary stats
summary.stats <- df.ribbon.comb %>%
  mutate(OOB_IQR = 1 -(Observed < round(q25, 2) | Observed > round(q75, 2)),
         OOB_95 = 1 -(Observed < round(q5, 2) | Observed > round(q95, 2)),
         Mix = as.numeric(Mix),
         # Filtered = (Mix == 5 | (Mix >=48 & Mix <= 67)),
         Abs.Error = abs(Error),
         Error_50 = ifelse(Observed >= q25 & Observed <= q75, 0,
                           ifelse(Observed < q25, q25 - Observed,
                                  q75 - Observed)),
         Error_95 = ifelse(Observed >= q5 & Observed <= q95, 0,
                           ifelse(Observed < q5, q5 - Observed,
                                  q95 - Observed)),
         Main = Observed > .5,
         Main50 = q75 > .5,
         Main95 = q95 > .5,
         Hit50 = ifelse(Main == T & Main50 == T, T, F),
         Hit95 = ifelse(Main == T & Main95 == T, T, F),
         Miss50 = ifelse(Main == T & Main50 == F, T, F),
         Miss95 = ifelse(Main == T & Main95 == F, T, F),
         FA50 = ifelse(Main == F & Main50 == T, T, F),
         FA95 = ifelse(Main == F & Main95 == T, T, F)) %>%
  group_by(Type, Source, Model) %>%
  summarise(Enc_50 = mean(OOB_IQR),
            Enc_95 = mean(OOB_95),
            W_IQR = mean(q75-q25),
            W_95 = mean(q95-q5),
            # R2 = cor(median, Observed),
            # ME = mean(Error),
            # MAE = mean(Abs.Error),
            # NSE = NSE(median, Observed),
            Mean.Obs = mean(Observed),
            MAE50 = mean(abs(Error_50)),
            MAE95 = mean(abs(Error_95)),
            ME50 = mean(Error_50),
            ME95 = mean(Error_95),
            NSE50 = 1 - (sum(Error_50^2)) / sum((Observed - mean(Observed))^2),
            NSE95 = 1 - (sum(Error_95^2)) / sum((Observed - mean(Observed))^2),
            CSI50 = sum(Hit50)/(sum(Hit50) + sum(Miss50) + sum(FA50)),
            CSI95 = sum(Hit95)/(sum(Hit95) + sum(Miss95) + sum(FA95)),
            HR50 = sum(Hit50) / (sum(Hit50) + sum(Miss50)),
            HR95 = sum(Hit95) / (sum(Hit95) + sum(Miss95)))

# Export stats
write.csv(summary.stats, "summary_stats.csv")

# Calculate CRPS
require(foreach)

# Temp dataframe 
temp <- rbind(MVN, MS) %>%
  mutate(Mix = as.factor(Mix),
         Type = as.factor(Type),
         Source = as.factor(Source), 
         Model = as.factor(Model))

# Loop crps through each model/mixture/source
crps.results <- foreach(i = 1:nlevels(temp$Mix), .combine = rbind) %:%
  foreach(k = 1:nlevels(temp$Source), .combine = rbind) %:%
  foreach(l = 1:nlevels(temp$Type), .combine = rbind) %:%
  foreach(m = 1:nlevels(temp$Model), .combine = rbind) %do% {
    
    mix <- levels(temp$Mix)[i] 
    source <- levels(temp$Source)[k] 
    type <- levels(temp$Type)[l] 
    model <- levels(temp$Model)[m]
    
    forecasts <- temp %>%
      filter(Mix == mix &
               Source == source &
               Type == type &
               Model == model)
    
    actuals <- df.am %>%
      mutate(Mix = as.factor(Mix)) %>%
      filter(Mix == mix &
               Source == source) 
    
    crps <- crps_sample(y = actuals$value, dat = forecasts$value)
    
    results <- data.frame(mix, source, type, model, crps, actuals$value)
  }

# Summarise CRPS results
crps.sum <- crps.results %>%
  group_by(type, model, source) %>%
  summarise(CRPS = mean(crps)) %>%
  rename(Model = model,
         Source = source,
         Type = type) %>%
  mutate(Source = ifelse(Source == "source_D", "Decontaminated",
                         ifelse(Source == "source_F", "Forest",
                                ifelse(Source == "source_PF", "Cropland", "Subsurface")))) %>%
  tbl_df() %>%
  mutate(Model = ifelse(Model == "MS", "MixSIAR", Model)) 

crps.sum.model <- crps.results %>%
  group_by(type, model) %>%
  summarise(CRPS = mean(crps)) %>%
  rename(Model = model,
         Type = type) %>%
  tbl_df() %>%
  mutate(Model = ifelse(Model == "MS", "MixSIAR", Model))

summary.stats1 <- summary.stats %>%
  rename(P50 = Enc_50,
         P95 = Enc_95,
         W50 = W_IQR,
         W95 = W_95) %>%
  mutate(Model = ifelse(Model == "MS", "MixSIAR", Model))

# Summarise results by model
summary.stats.model <- df.ribbon.comb %>%
  mutate(OOB_IQR = 1 -(Observed < round(q25, 2) | Observed > round(q75, 2)),
         OOB_95 = 1 -(Observed < round(q5, 2) | Observed > round(q95, 2)),
         Mix = as.numeric(Mix),
         Abs.Error = abs(Error),
         Error_50 = ifelse(Observed >= q25 & Observed <= q75, 0,
                           ifelse(Observed < q25, q25 - Observed,
                                  q75 - Observed)),
         Error_95 = ifelse(Observed >= q5 & Observed <= q95, 0,
                           ifelse(Observed < q5, q5 - Observed,
                                  q95 - Observed)),
         Main = Observed > .5,
         Main50 = q75 > .5,
         Main95 = q95 > .5,
         Hit50 = ifelse(Main == T & Main50 ==T, T, F),
         Hit95 = ifelse(Main == T & Main95 ==T, T, F),
         Miss50 = ifelse(Main == T & Main50 == F, T, F),
         Miss95 = ifelse(Main == T & Main95 == F, T, F),
         FA50 = ifelse(Main == F & Main50 == T, T, F),
         FA95 = ifelse(Main == F & Main95 == T, T, F)) %>%
  group_by(Type, Model) %>%
  summarise(Enc_50 = mean(OOB_IQR),
            Enc_95 = mean(OOB_95),
            W_IQR = mean(q75-q25),
            W_95 = mean(q95-q5),
            # R2 = cor(median, Observed),
            # ME = mean(Error),
            # MAE = mean(Abs.Error),
            # NSE = NSE(median, Observed),
            Mean.Obs = mean(Observed),
            MAE50 = mean(abs(Error_50)),
            MAE95 = mean(abs(Error_95)),
            ME50 = mean(Error_50),
            ME95 = mean(Error_95),
            NSE50 = 1 - (sum(Error_50^2)) / sum((Observed - mean(Observed))^2),
            NSE95 = 1 - (sum(Error_95^2)) / sum((Observed - mean(Observed))^2),
            CSI50 = sum(Hit50)/(sum(Hit50) + sum(Miss50) + sum(FA50)),
            CSI95 = sum(Hit95)/(sum(Hit95) + sum(Miss95) + sum(FA95)),
            HR50 = sum(Hit50) / (sum(Hit50) + sum(Miss50)),
            HR95 = sum(Hit95) / (sum(Hit95) + sum(Miss95))) %>%
  rename(P50 = Enc_50,
         P95 = Enc_95,
         W50 = W_IQR,
         W95 = W_95)

# Export summary 
write.csv(summary.stats1 %>%
            inner_join(crps.sum), "summary_stats_SMT.csv")
write.csv(summary.stats.model %>%
            inner_join(crps.sum.model), "summary_stats_MT.csv")

# Plot CRPS results
ggplot(crps.results %>%
         mutate(model = ifelse(model == "MS", "MixSIAR", model)) %>%
         mutate(source = ifelse(source == "source_D", "Decontaminated",
                                ifelse(source == "source_F", "Forest",
                                       ifelse(source == "source_PF", "Cropland", "Subsurface")))) %>%
         mutate(type = ifelse(type == "LabMix", "Laboratory", "Mathematical")), 
       aes(x = actuals.value, y = crps)) +
  geom_point(aes(colour = model, shape = model)) +
  geom_smooth(aes(colour = model), se = F) +
  facet_grid(rows = vars(type), cols = vars(source)) +
  theme_bw() +
  scale_shape_manual(values = c(1,2)) +
  scale_color_viridis(discrete = T, begin = .8, end = 0.2, option = "D") +
  xlim(0,1) +
  ylim(0,1) +
  xlab("\nSource proportion") +
  ylab("CRPS\n ") +
  theme(panel.spacing = unit(1, "lines"),
        legend.position = "top",
        legend.title = element_blank(),
        text = element_text(size = 12),
        axis.text.x = element_text(angle = 45, hjust = 1)) 

# Export  
ggsave("obs_crps_final.png", width = 10, height = 6)

#density plots for mixture 10
# Combine
df.dens <- rbind(MS, MVN) %>%
  mutate(Source = ifelse(Source == "source_D", "Decontaminated",
                         ifelse(Source == "source_F", "Forest",
                                ifelse(Source == "source_PF", "Cropland", "Subsurface"))),
         Model = ifelse(Model == "MS", "MixSIAR", Model)) %>%
  filter(Mix == "10")


ggplot(df.dens, aes(x = value, y = ..scaled..)) +
  geom_density(aes(colour = Source, fill = Source), alpha = 0.2) +
  # geom_vline(data = df.obs %>%
  #              filter(Mix == "10"), 
  #            aes(xintercept = Observed, colour = Source),
  #            linetype = "dashed") +
  scale_color_viridis_d() +
  scale_fill_viridis_d() +
  facet_grid(rows = vars(Model), cols = vars(Type)) +
  theme_bw() +
  ylab("Scaled density\n") +
  xlab("\nSource proportion") +
  theme(legend.title = element_blank(),
        legend.position = "top",
        panel.spacing = unit(1, "lines"))

ggsave("Denstiy_Mix10_Final.png", width = 5, height = 5)

# Scatter plot Lab vs Math
df.scatter <- df.comb %>%
  mutate(Model = ifelse(Model == "MS", "MixSIAR", "BMM"),
         Quant = ifelse(Quant == "q5", ".05 quantile", 
                 ifelse(Quant == "q25", ".25 quantile",
                 ifelse(Quant == "median", ".50 quantile",
                 ifelse(Quant == "q75", ".75 quantile",
                 ifelse(Quant == "q95", ".95 quantile",
                 Quant)))))) %>%
  dcast(Source + Mix + Model + Quant ~ Type, value.var = "value") %>%
  tbl_df()

dat_text_MS <- df.scatter %>%
  group_by(Model, Quant) %>%
  summarise(r = round(cor(LabMix, MathMix), 2)) %>%
  filter(Model == "MixSIAR") %>%
  mutate(Label = paste0("MixSIAR r = ", r))

dat_text_MVN <- df.scatter %>%
  group_by(Model, Quant) %>%
  summarise(r = round(cor(LabMix, MathMix), 2)) %>%
  filter(Model == "BMM") %>%
  mutate(Label = paste0("BMM r = ", r))
  
ggplot(df.scatter, aes(x = LabMix, y = MathMix)) +
  geom_point(aes(colour = Model, shape = Model)) +
  theme_bw() +
  facet_wrap(~Quant) +
  scale_shape_manual(values = c(1,2)) +
  scale_color_viridis(discrete = T, begin = .8, end = 0.2, option = "D") +
  geom_abline(intercept = 0, slope = 1) +
  xlim(0,1) +
  ylim(0,1) +
  xlab("\nLaboratory Mixtures") +
  ylab("Mathematical Mixtures\n ") +
  theme(panel.spacing = unit(1, "lines"),
        legend.position = "top",
        legend.title = element_blank()) +
  geom_text(
    data    = dat_text_MS,
    mapping = aes(x = 0.23, y = 0.95, label = Label),
    # hjust   = -0.1,
    # vjust   = -1,
    size = 3.5,
    colour = viridis(n = 1, begin = 0.2)) +
  geom_text(
    data    = dat_text_MVN,
    mapping = aes(x = 0.18, y = 0.82, label = Label),
    # hjust   = -0.1,
    # vjust   = -1,
    size = 3.5,
    colour = viridis(n = 1, begin = 0.8))


ggsave("Lab_Math_Cor_Final.png", width = 8, height = 6.5)







