#
#  R Code for reproducing results of the article
#
#  "Mixture density networks for the indirect estimation of reference intervals"
#
#  T.Hepp et al.
#
########################################### 

### Custom utility functions

tanh.base <- function(x,number.of.bases){
  
  b <- -number.of.bases*(min(x)+(1:number.of.bases-1)/(number.of.bases-1)*diff(range(x)))/diff(range(x))
  w <- number.of.bases/diff(range(x))
  
  out <- tanh(cbind(1,x)%*%rbind(b,w))
  attr(out,"weights") <- rbind(b,w)
  
  return(out)
}

glorot_uniform <- function(dim){
  
  out <- MDN.vectomat(rep(0,dim$total),dim)
  
  out$xh[-1,] <- runif(length(out$xh[-1,]),-sqrt(6)/sqrt((dim$xh[1]-1)+dim$xh[2]),sqrt(6)/sqrt((dim$xh[1]-1)+dim$xh[2]))
  
  out$ha[-1,] <- runif(length(out$hm[-1,]),-sqrt(6)/sqrt((dim$hm[1]-1)+dim$hm[2]*3),sqrt(6)/sqrt((dim$hm[1]-1)+dim$hm[2]*3))
  out$hm[-1,] <- runif(length(out$hm[-1,]),-sqrt(6)/sqrt((dim$hm[1]-1)+dim$hm[2]*3),sqrt(6)/sqrt((dim$hm[1]-1)+dim$hm[2]*3))
  out$hs[-1,] <- runif(length(out$hm[-1,]),-sqrt(6)/sqrt((dim$hm[1]-1)+dim$hm[2]*3),sqrt(6)/sqrt((dim$hm[1]-1)+dim$hm[2]*3))
  
  return(out)
}


predict_and_rescale <- function(run,model,x,param="mu",comp=NULL){
  
  rescale_x <- (x-run$rescale$mean.x)/run$rescale$sd.x
  
  fit <-  MDN.predict(run[[model]],rescale_x,param,comp)
  
  if(param=="mu")
    fit <- fit*run$rescale$sd.y+run$rescale$mean.y
  
  if(param=="sigma"){
    if(grepl("ADAM",model)){
      fit <- log(1+fit)
    }
    fit <- fit*run$rescale$sd.y
  }
  return(fit)
}



sortstuff <- function(simlist,par="mu",x=seq(0,1,l=100)){
  
  for(s in 1:length(simlist)){
    
    corder_em <- order(colMeans(predict_and_rescale(simlist[[s]],"EM",x,par)))
    if(par=="alpha") corder_em <- rev(corder_em)
    simlist[[s]]$EM$W$ha <- matrix(simlist[[s]]$EM$W$ha[,corder_em],1) # different shape for alpha weights
    simlist[[s]]$EM$W$hm <- simlist[[s]]$EM$W$hm[,corder_em]
    simlist[[s]]$EM$W$hs <- simlist[[s]]$EM$W$hs[,corder_em]
    
    corder_bfgs_acons <- order(colMeans(predict_and_rescale(simlist[[s]],"BFGS_acons",x,par)))
    if(par=="alpha") corder_bfgs_acons <- rev(corder_bfgs_acons)
    simlist[[s]]$BFGS_acons$W$ha <- matrix(simlist[[s]]$BFGS_acons$W$ha[,corder_bfgs_acons],1) # different shape for alpha weights
    simlist[[s]]$BFGS_acons$W$hm <- simlist[[s]]$BFGS_acons$W$hm[,corder_bfgs_acons]
    simlist[[s]]$BFGS_acons$W$hs <- simlist[[s]]$BFGS_acons$W$hs[,corder_bfgs_acons]
    
    corder_bfgs_cinit_acons <- order(colMeans(predict_and_rescale(simlist[[s]],"BFGS_cinit_acons",x,par)))
    if(par=="alpha") corder_bfgs_cinit_acons <- rev(corder_bfgs_cinit_acons)
    simlist[[s]]$BFGS_cinit_acons$W$ha <- matrix(simlist[[s]]$BFGS_cinit_acons$W$ha[,corder_bfgs_cinit_acons],1) # different shape for alpha weights
    simlist[[s]]$BFGS_cinit_acons$W$hm <- simlist[[s]]$BFGS_cinit_acons$W$hm[,corder_bfgs_cinit_acons]
    simlist[[s]]$BFGS_cinit_acons$W$hs <- simlist[[s]]$BFGS_cinit_acons$W$hs[,corder_bfgs_cinit_acons]
    
    corder_bfgs <- order(colMeans(predict_and_rescale(simlist[[s]],"BFGS",x,par)))
    if(par=="alpha") corder_bfgs <- rev(corder_bfgs)
    simlist[[s]]$BFGS$W$ha <- simlist[[s]]$BFGS$W$ha[,corder_bfgs]
    simlist[[s]]$BFGS$W$hm <- simlist[[s]]$BFGS$W$hm[,corder_bfgs]
    simlist[[s]]$BFGS$W$hs <- simlist[[s]]$BFGS$W$hs[,corder_bfgs]
    
    corder_bfgs_custom <- order(colMeans(predict_and_rescale(simlist[[s]],"BFGS_custom",x,par)))
    if(par=="alpha") corder_bfgs_custom <- rev(corder_bfgs_custom)
    simlist[[s]]$BFGS_custom$W$ha <- simlist[[s]]$BFGS_custom$W$ha[,corder_bfgs_custom]
    simlist[[s]]$BFGS_custom$W$hm <- simlist[[s]]$BFGS_custom$W$hm[,corder_bfgs_custom]
    simlist[[s]]$BFGS_custom$W$hs <- simlist[[s]]$BFGS_custom$W$hs[,corder_bfgs_custom]
    
    corder_adam <- order(colMeans(predict_and_rescale(simlist[[s]],"ADAM",x,par)))
    if(par=="alpha") corder_adam <- rev(corder_adam)
    simlist[[s]]$ADAM$W$ha <- simlist[[s]]$ADAM$W$ha[,corder_adam]
    simlist[[s]]$ADAM$W$hm <- simlist[[s]]$ADAM$W$hm[,corder_adam]
    simlist[[s]]$ADAM$W$hs <- simlist[[s]]$ADAM$W$hs[,corder_adam]
    
    corder_adam_custom <- order(colMeans(predict_and_rescale(simlist[[s]],"ADAM_custom",x,par)))
    if(par=="alpha") corder_adam_custom <- rev(corder_adam_custom)
    simlist[[s]]$ADAM_custom$W$ha <- simlist[[s]]$ADAM_custom$W$ha[,corder_adam_custom]
    simlist[[s]]$ADAM_custom$W$hm <- simlist[[s]]$ADAM_custom$W$hm[,corder_adam_custom]
    simlist[[s]]$ADAM_custom$W$hs <- simlist[[s]]$ADAM_custom$W$hs[,corder_adam_custom]
  }
  
  return(simlist)
}
