###################################################################################################################
# Function to fit bivariate Poisson-lognormal species-abundance distributions by MLE method (Engen et al. 2002)
# input SADdat: Community matrix or Species-abundance distribution over time (make sure the time index is matching up the columns)
# input reportyears: Reported years
# input reefname : Reef names
# input ltmp_transectlarge_sprowindex : speices index from the large transect
# output MLout: MLEs of bivariate Poission-lognormal species-abundance distributions
###################################################################################################################
ltmpPoiLogFit_fn <- function(SADdat,reportyears,reefname,ltmp_transectlarge_sprowindex)
{
  library(poilog)
  library(tidyverse)
  library(optimx)
  idx_small <-  which(!(1:232%in%ltmp_transectlarge_sprowindex)) #small transect species
  idx_large <- ltmp_transectlarge_sprowindex #large transect species
  nlag <- length(reportyears)
  tmp <- expand.grid(m=reportyears,n=reportyears)
  tmp <- tmp%>%mutate(idx=m>n)%>%filter(idx==TRUE)%>%mutate(tlag=m-n)
  parmsall <- c()
  for(i in 1:nrow(tmp)){
    m <- which(reportyears%in%tmp[i,1])
    n <- which(reportyears%in%tmp[i,2])
    xy <- cbind(SADdat[,n],SADdat[,m])
    pair_data <- xy[idx_small,]
    pair_data2 <- xy[idx_large,]
    pair_data <- pair_data[rowSums(pair_data)!=0,]
    pair_data2 <- pair_data2[rowSums(pair_data2)!=0,]
    nsp.x <- which(rbind(pair_data,pair_data2)[,1]!=0) %>% length
    nsp.y <- which(rbind(pair_data,pair_data2)[,2]!=0) %>% length
    est <- try(bipoilog_mle(pair_data=pair_data,pair_data2=pair_data2,nu=c(1,0.2),start_values=c(1,1,2,2,0),method="Nelder-Mead"))
    if(!is.na(est[1]))
    {
      p1 <- 1-dpoilog(0,est$mu1,est$sig1)
      p2 <- 1-dpoilog(0,est$mu2,est$sig2)
      parms <- c(est%>%unlist,strue.x=nsp.x/p1,strue.y=nsp.y/p2)
    }
    else
    {
      est2 <- try(bipoilog_mle(pair_data=pair_data,pair_data2=pair_data2,nu=c(1,0.2),start_values=c(1,1,2,2,0),method="nlminb"))
      if(!is.na(est2[1]))
      {
        p1 <- 1-dpoilog(0,est2$mu1,est2$sig1)
        p2 <- 1-dpoilog(0,est2$mu2,est2$sig2)
        parms <- c(est2%>%unlist,strue.x=nsp.x/p1,strue.y=nsp.y/p2)
      }
      else
      {
        est3 <- try(bipoilog_mle(pair_data=pair_data,pair_data2=pair_data2,nu=c(1,0.2),start_values=c(1,1,2,2,0),method="Rcgmin"))
        if(!is.na(est3[1]))
        {
          p1 <- 1-dpoilog(0,est3$mu1,est3$sig1)
          p2 <- 1-dpoilog(0,est3$mu2,est3$sig2)
          parms <- c(est3%>%unlist,strue.x=nsp.x/p1,strue.y=nsp.y/p2)
        }
        else {
          parms <- NA 
        }
      }
    }
    parmsall <- rbind(parmsall,parms) 
  }
  out <- bind_cols(tmp,parmsall,row.names=NULL)
  out <- out %>% mutate(reefname=reefname)
  return(list(out))
}
####################################################################################################################
# The function bipoilogMLE is inflexible in this application and that is why
# it is better to work with a user-defined function. Problems with the
# bipoilogMLE function:
# 1. It doesn't give convergence information which is extremely important when
# there are a lot of zeroes in the dataset.
# 2. It doesn't have many optimization methods available. The function optimx()
# provides more methods.
# 3. It can't be modified to accept variable sampling effort as present in this
# application.
# This script provides a flexible function to fit a bivariate Poisson Lognormal
# distribution.

### Transformation Functions. ----
# These are the functions to transform the rho parameter for optimization.
imlogit <- function(n) {
  val <- (exp(n) - 1)/(exp(n) + 1)
  return(val)
}
# Modified logit transform.
mlogit <- function(n) {
  return(log(n+1) - log(1-n))
}
###  Wrapper functions for dbipoilog - Argument Vectorization ----
vdbipoilog <- function(n1,
                       n2,
                       pars) {
  return(dbipoilog(n1,n2,pars[1],pars[2],pars[3],pars[4],pars[5]))
}
###  Function to calculate the biPLN log-likelihood ----
bipoilog_zt_log_likelihood <- function(pars,
                                       pair_data,
                                       pair_data2=NULL,
                                       nu = c(1, 1)) {
  ll <- 0
  
  if( is.null(pair_data2) ) {
    # Distribution parameters applying the necessary links.
    # Ordered as: mu1,mu2,sig1,sig2,rho.
    pars <- c(pars[1] + log(nu[1]), pars[2] + log(nu[2]),
              exp(pars[3]), exp(pars[4]),
              imlogit(pars[5]))
    
    ll <- sum(pair_data$n * log(vdbipoilog(pair_data[[1]],pair_data[[2]],pars))) -
      sum(pair_data$n) * log(1-vdbipoilog(0,0,pars))
  } else {
    pars1 <- c(pars[1] + log(nu[1]), pars[2] + log(nu[1]),
               exp(pars[3]), exp(pars[4]),
               imlogit(pars[5]))
    
    pars2 <- c(pars[1] + log(nu[2]), pars[2] + log(nu[2]),
               exp(pars[3]), exp(pars[4]),
               imlogit(pars[5]))
    
    ll <- sum(pair_data$n * log(vdbipoilog(pair_data[[1]],pair_data[[2]],pars1))) -
      sum(pair_data$n) * log(1-vdbipoilog(0,0,pars1))
    
    ll <- ll + sum(pair_data2$n * log(vdbipoilog(pair_data2[[1]],pair_data2[[2]],pars2))) -
      sum(pair_data2$n) * log(1-vdbipoilog(0,0,pars2))
  }
  
  # Returning biPLN likelihood. Zero-truncated.
  return(ll)
}
###  Function to calculate the biPLN MLE ---
# If the dataset contains two different sampling efforts, split the dataset into 
# two parts. Pass one dataset to pair_data, the second to pair_data2, and also pass
# the corresponding nu values to the nu argument.
bipoilog_mle <- function(pair_data,
                         pair_data2=NULL,
                         nu = c(1, 1),
                         start_values=c(1, 1, 2, 2, 0),
                         method="Nelder-Mead") {
  # Initial parameter estimates.
  init_pars <- c(start_values[1],start_values[2],log(start_values[3]),log(start_values[4]),mlogit(start_values[5]))
  
  # init_pars <- c(1,1,log(2),log(2),mlogit(0.5))
  names(init_pars) <- c("mu1","mu2","sig1","sig2","rho")
  
  # Tabulating pairs.
  pair_data <- pair_data[rowSums(pair_data)>0, ]
  pair_data <- as.data.frame(pair_data) %>% group_by_all %>% count
  
  if( !is.null(pair_data2) ) {
    pair_data2 <- pair_data2[rowSums(pair_data2)>0, ]
    pair_data2 <- as.data.frame(pair_data2) %>% group_by_all %>% count
  }
  
  # Optimization of the likelihood function.
  mle_fit <- optimx(init_pars, bipoilog_zt_log_likelihood,
                    method = method,
                    itnmax = 10000,
                    pair_data = pair_data,
                    pair_data2 = pair_data2,
                    nu = nu,
                    control = list(kkt=FALSE, maximize=TRUE, maxit=1000))
  
  # Formatted model fit output.
  mle_fit <- data.frame(mu1 = mle_fit[1,1],
                        mu2 = mle_fit[1,2],
                        sig1 = exp(mle_fit[1,3]),
                        sig2 = exp(mle_fit[1,4]),
                        rho = imlogit(mle_fit[1,5]),
                        ll = mle_fit$value,
                        conv = mle_fit$convcode,
                        method = method)
  
  return(mle_fit)
}
####################################################################################################################


####################################################################################################################
# loading packages required for MLE of bivariate poisson-lognormal fits
library(poilog)
library(optimx)
library(tidyverse)
# loading package required for parallel computing
library(foreach)
library(doParallel)
cl <- makeCluster(4)
registerDoParallel(cl)
getDoParWorkers()
#
load("dsrltmp2024_meta.RData")
# obtain MLE results of Poisson-lognormal (PLN) fits using "ltmpPoiLogFit_fn"
test_poilog_mle_out <- foreach(i=1:length(gbrfishes$ltmp_fishcomm_list),.combine=c) %dopar% {
  out <- ltmpPoiLogFit_fn(SADdat=gbrfishes$ltmp_fishcomm_list[[i]],
                          reportyears=ltmp_reportyears,
                          reefname=ltmp_reefnames[i],
                          ltmp_transectlarge_sprowindex=ltmp_transectlarge_sprowindex)
}
Sys.time()
stopCluster(cl)
