#
#  R Code for reproducing results of the article
#
#  "Mixture density networks for the indirect estimation of reference intervals"
#
#  T.Hepp et al.
#
########################################### 

### EM-algorithm

# WARNING: Problem with PMAT if dep var is not named "y"
# QUICKFIX: Use first formula object and grab name of dep var

require(gamlss)

lcdr <- function(formula,data=NULL,init=NULL,threshold=0.001,tstart=NULL,
                 run=NULL){ #cluster only
  
  tst <- tstart
  
  data <- data.frame(data)
  M <- length(formula)/2
  
  # initial weights
  if(!is.matrix(init)){
    stop(cat("Please provide weights for initialization in matrix format"))
  }
  
  # models
  MODS <- lapply(1:M,function(m){
    gamlss(formula = eval(parse(text=formula[2*m-1])),
           sigma.formula = eval(parse(text=formula[2*m])),
           data=data.frame(data),weights=init[,m],trace=F)
  })
  
  # mixt. weights
  A <- apply(init,2,sum)/nrow(data)
  
  # cond. pdf of each comp.
  PMAT <- sapply(1:M,function(m){
    dnorm(with(data,get(all.vars(formula[[1]])[1])),MODS[[m]]$mu.fv,MODS[[m]]$sigma.fv)*A[m]
  })
  
  # New model weights
  WMAT <- PMAT/rowSums(PMAT)
  
  lL <- c(-Inf,sum(log(rowSums(PMAT))))
  lLvec <- lL[2]
  
  error <- tryCatch({
    while((diff(lL))>threshold){
      
      # M-Step
      
      # models
      MODS <- lapply(1:M,function(m){
        gamlss(formula = eval(parse(text=formula[2*m-1])),
               sigma.formula = eval(parse(text=formula[2*m])),
               data=data.frame(data),weights=WMAT[,m],start.from=MODS[[m]],trace=F)
      })
      
      # mixt. weights
      A <- apply(WMAT,2,sum)/nrow(data)
      
      # E-Step
      
      # cond. pdf of each comp.
      PMAT <- sapply(1:M,function(m){
        dnorm(with(data,get(all.vars(formula[[1]])[1])),MODS[[m]]$mu.fv,MODS[[m]]$sigma.fv)*A[m]
      })
      
      # New model weights
      WMAT <- PMAT/rowSums(PMAT)
      
      lL[1] <- lL[2]
      lL[2] <- sum(log(rowSums(PMAT)))
      
      lLvec <- c(lLvec,lL[2])
    }
  },error=function(e){
    cat("ID",run,":",conditionMessage(e), "\n")
    return(conditionMessage(e))
  })
  
  FITARR <- array(NA,c(nrow(data),2,M))
  FITARR[,1,] <- sapply(1:M,function(m){MODS[[m]]$mu.fv})
  FITARR[,2,] <- sapply(1:M,function(m){MODS[[m]]$sigma.fv})
  
  dimnames(FITARR)[[2]] <- c("mu","sigma")
  dimnames(FITARR)[[3]] <- paste0("Comp",1:M)
  
  list(mucoefs=lapply(1:M,function(m){coef(MODS[[m]],what="mu")}),
       sicoefs=lapply(1:M,function(m){coef(MODS[[m]],what="sigma")}),
       fit=FITARR,lL=lLvec,alpha=A,data=data,error=error)
  
}
