.packageName <- "ouch"
# This file is part of the OUCH package.
# Author: Aaron A. King <king at tiem dot utk dot edu>
# It is distributed under the GNU Public License (see the file GPL
# included)
# The OUCH package is maintained at
#          http://www.tiem.utk.edu/~king/ouch/
#
brown.fit <- function (data, node, ancestor, times) {
  pt <- parse.tree(node,ancestor,times)
  n <- pt$N
  v <- pt$branch.times
  w <- matrix(data=1,nrow=pt$N,ncol=1)
  dat <- data[pt$term]
  no.dats <- which(is.na(dat))
  if (length(no.dats) > 0)
    stop("Missing data on terminal nodes: ",node[pt$term[no.dats]])
  g <- glssoln(w,dat,v)
  theta <- g$coeff
  e <- g$residuals
  sigma <- sqrt((e %*% solve(v,e))/n)
  dim(sigma) <- 1
  u = n * (1 + log(2*pi*sigma*sigma)) + log(det(v))
  dim(u) <- 1
  df <- 2
  list(sigma=sigma,theta.0=theta,u=u,aic=u+2*df,sic=u+log(n)*df,df=df)
}

brown.dev <- function(n = 1, node, ancestor, times, sigma, theta) {
  pt <- parse.tree(node,ancestor,times)
  v <- pt$branch.times
  x <- rmvnorm(n, rep(theta,dim(v)[1]), as.numeric(sigma^2)*v)
  data.frame(x)
}


glssoln <- function(a, x, v, tol = 1e-12) {
  n <- length(x);
  vh <- t(chol(v));
  s <- svd(forwardsolve(vh,a));
                                        #   Can we be certain that the singular values are sorted in
                                        #   decreasing order?  (Probably not)
                                        #   k <- order(s$d,decreasing=T);
  svals <- s$d[s$d > tol * max(s$d)];
  r <- length(svals);
  svals <-  diag(1/svals,nrow=r,ncol=r);
  y <- (s$v[,1:r] %*% (svals %*% t(s$u[,1:r]))) %*% forwardsolve(vh,x);
  e <- a %*% y - x;
  dim(y) <- dim(y)[1];
  dim(e) <- n;
  return(list(coeff=y,residuals=e));
}

hansen.fit <- function (data, node, ancestor, times, regimes,
                        guess=0, interval=c(0.001,20), tol=1e-12) {
  pt <- parse.tree(node,ancestor,times,regimes);
  n <- pt$N;
  dat <- data[pt$term]
  no.dats <- which(is.na(dat))
  if (length(no.dats) > 0)
    stop("Missing data on terminal nodes: ",node[pt$term[no.dats]])
  r <- optimize(badness,interval=log(interval),
                lower=log(interval[1]),upper=log(interval[2]),
                tol=tol,maximum=F,dat,pt);
  alpha = exp(r$minimum);

  w <- weight.matrix(alpha, pt);
  v <- scaled.covariance.matrix(alpha, pt);
  g <- glssoln(w,dat,v);
  theta <- g$coeff;
  names(theta) <- paste('theta',c('0',as.character(pt$regime.set)),sep='.')
  e <- g$residuals;
  sigma <- sqrt((e %*% solve(v,e))/n);
  dim(sigma) <- 1;
  u = r$objective;
  dim(u) <- 1;
  df <- pt$R+3;
  return(list(alpha=alpha,sigma=sigma,theta=theta,u=u,aic=u+2*df,sic=u+log(n)*df,df=df));
}

hansen.dev <- function(n = 1, node, ancestor, times, regimes, alpha, sigma, theta) {
  pt <- parse.tree(node,ancestor,times,regimes);
  w <- weight.matrix(alpha, pt);
  v <- scaled.covariance.matrix(alpha, pt);
  x <- rmvnorm(n, as.vector(w %*% theta), as.numeric(sigma^2)*v);
  return(data.frame(x))
}

badness <- function (alpha, data, parsed.tree) {
  a <- exp(alpha);
  n <- length(data);
  w <- weight.matrix(a, parsed.tree);
  v <- scaled.covariance.matrix(a, parsed.tree);
  g <- glssoln(w,data,v);
  e <- g$residuals;
  sigmasq <- (e %*% solve(v,e)) / n;
  dim(sigmasq) <- 1;
  u <- n * (1 + log(2*pi*sigmasq)) + log(det(v));
  dim(u) <- 1;
  return(u);
}

weight.matrix <- function (alpha, parsed.tree) {
  N <- parsed.tree$N;
  R <- parsed.tree$R;
  tree.depth <- parsed.tree$tree.depth;
  ep <- parsed.tree$epochs;
  beta <- parsed.tree$beta;
  W <- matrix(data=0,nrow=N,ncol=R+1);
  W[,1] <- exp(-alpha*tree.depth);      
  for (i in 1:N) {
    delta <- diff(exp(alpha*(ep[[i]]-tree.depth)));
    for (k in 1:R) {
      W[i,k+1] <- -sum(delta * beta[[i+N*(k-1)]]);
    }
  }
  return(W);
}

scaled.covariance.matrix <- function (alpha, parsed.tree) {
  tree.depth <- parsed.tree$tree.depth;
  bt <- parsed.tree$branch.times;			 
  if (alpha == 0) {
    V <- bt;
  } else {
    a <- 2*alpha;
    V <- exp(-a*tree.depth) * expm1(a*bt) / a;
  }
}

parse.tree <- function (nodenames, ancestors, times, regime.specs=NULL) {
  nodenames <- as.character(nodenames)
  ancestors <- as.character(ancestors)
  if (!is.valid.ouch.tree(nodenames,ancestors,times,regime.specs))
    stop('the specified tree is not in valid ouch format')
  term <- terminal.twigs(nodenames,ancestors) # get rownumbers of terminal nodes
  N <- length(term)                     # number of terminal nodes
  anc <- ancestor.numbers(nodenames,ancestors)
  bt <- branch.times(anc,times,term)   # absolute times of branch points
  e <- epochs(anc,times,term)
  if (is.null(regime.specs)) {          # useful for BM models
    pt <- list(
               N=N,
               tree.depth = max(times),
               term=term,
               branch.times=bt,
               epochs=e
               )
  } else {                              # useful for Hansen models
    reg <- set.of.regimes(anc,as.factor(regime.specs))
    pt <- list(
               N=N,
               R=length(reg),
               tree.depth = max(times),
               term=term,
               branch.times=bt,
               epochs=e,
               regime.set=reg,
               beta=regimes(anc,times,as.factor(regime.specs),term)
               )
  }
  return(pt)
}

ancestor.numbers <- function (nodenames, ancestors) { # map ancestor names to row numbers
  sapply(ancestors,function(x)charmatch(x,nodenames),USE.NAMES=F)
}

terminal.twigs <- function (nodenames, ancestors) { # numbers of terminal nodes
  which(nodenames %in% setdiff(nodenames,unique(ancestors)))
}

branch.times <- function (ancestors, times, term) {

  N <- length(term)
  tree.depth <- max(times)              # it is assumed that the root node is at time=0

  bt <- matrix(data=0,nrow=N,ncol=N)

  bt[1,1] <- tree.depth
  for (i in 2:N) {
    pedi <- pedigree(ancestors,term[i])
    for (j in 1:(i-1)) {
      pedj <- pedigree(ancestors,term[j])
      for (k in 1:length(pedi)) {
        if (any(pedj == pedi[k])) break
      }
      bt[j,i] <- bt[i,j] <- times[pedi[k]]
    }
    bt[i,i] <- tree.depth
  }
  bt
}

epochs <- function (ancestors, times, term) {
  N <- length(term)
  e <- vector(length=N,mode="list")
  for (k in 1:N) {
    p <- pedigree(ancestors,term[k])
    e[[k]] <- times[p]	
  }
  e
}

set.of.regimes <- function (ancestors, regime.specs) {
  unique(regime.specs[!is.root.node(ancestors)])
}

regimes <- function (ancestors, times, regime.specs, term) {
  N <- length(term)
  reg <- set.of.regimes(ancestors,regime.specs)
  R <- length(reg)
  beta <- vector(R*N, mode="list")
  for (i in 1:N) {
    for (k in 1:R) {
      p <- pedigree(ancestors, term[i])
      n <- length(p)
      beta[[i + N*(k-1)]] <- as.integer(regime.specs[p[1:(n-1)]] == reg[k])
    }
  }    
  beta
}

pedigree <- function (anc, k) {
  p <- k
  k <- anc[k]
  while (!is.root.node(k)) {
    if (k %in% p) stop('this is no tree: circularity detected at node ', k)
    p <- c(p,k)
    k <- anc[k]
  }
  p
}

is.root.node <- function (anc) {
  is.na(anc)
}

rmvnorm <- function (n = 1, mu, sigma, tol = 1e-06) {
  p <- length(mu);
  if (!all(dim(sigma) == c(p,p)))
    stop("incompatible arguments");
  cf <- chol(sigma,pivot=F);
  X <- matrix(mu,n,p,byrow=T) + matrix(rnorm(p*n),n) %*% cf;
  if (n == 1) {
    return(drop(X));
  }  else {
    return(X);
  }
}

tree.plot <- function (node, ancestor, times, names = NULL, regimes = NULL) {

  node <- as.character(node)
  ancestor <- as.character(ancestor)
  if (!is.valid.ouch.tree(node,ancestor,times,regimes))
    stop("the tree is not in valid ouch format");

  rx <- range(times,na.rm=T)
  rxd <- 0.1*diff(rx)

  anc <- ancestor.numbers(node,ancestor)

  if (is.null(regimes))
    regimes <- factor(rep(1,length(anc)))

  levs <- levels(as.factor(regimes))
  palette <- rainbow(length(levs))

  for (r in 1:length(levs)) {
    y <- tree.layout(anc)
    x <- times
    f <- which(!is.root.node(anc) & regimes == levs[r])
    pp <- anc[f]
    X <- array(data=c(x[f], x[pp], rep(NA,length(f))),dim=c(length(f),3))
    Y <- array(data=c(y[f], y[pp], rep(NA,length(f))),dim=c(length(f),3))
    oz <- array(data=1,dim=c(2,1))
    X <- kronecker(t(X),oz)
    Y <- kronecker(t(Y),oz)
    X <- X[2:length(X)]
    Y <- Y[1:(length(Y)-1)]
    C <- rep(palette[r],length(X))
    if (r > 1) par(new=T)
    par(yaxt='n')
    plot(X,Y,type='l',col=C,xlab='time',ylab='',xlim = rx + c(-rxd,rxd),ylim=c(0,1))
    if (!is.null(names))
      text(X[seq(1,length(X),6)],Y[seq(1,length(Y),6)],names[f],pos=4)
  }
}

tree.layout <- function (anc) {
  root <- which(is.root.node(anc))
  arrange.tree(root,anc)
}

arrange.tree <- function (root, anc) {
  k <- which(anc==root)
  n <- length(k)
  reltree <- rep(0,length(anc))
  reltree[root] <- 0.5
  p <- list()
  if (n > 0) {
    m <- rep(0,n)
    for (j in 1:n) {
      p[[j]] <- arrange.tree(k[j],anc)
      m[j] <- length(which(p[[j]] != 0))
    }
    cm <- c(0,cumsum(m))
    for (j in 1:n) {
      reltree <- reltree + (cm[j]/sum(m))*(p[[j]] != 0) + (m[j]/sum(m))*p[[j]]
    }
  }
  reltree
}


is.valid.ouch.tree <- function (node, ancestor, times, regimes=NULL) {
  valid <- TRUE
  node <- as.character(node)
  ancestor <- as.character(ancestor)
  n <- length(node)
  if (length(unique(node)) != n) {
    warning('node names must be unique')
    valid <- FALSE
  }
  if (
      (length(ancestor) != n) ||
      (length(times) != n)
      ) {
    warning('invalid tree: all columns must be of the same length')
    valid <- FALSE
  }
  if (!is.null(regimes) && (length(regimes) != n)) {
    warning('regimes must be the same length as the other columns')
    valid <- FALSE
  }
  root <- which(is.root.node(ancestor))
  if (length(root) != 1) {
    warning('the tree must have a unique root node, designated by its having ancestor = NA')
    return(FALSE)
  }
  term <- as.list(terminal.twigs(node,ancestor))
  if (length(term) <= 0) {
    warning("there ought to be at least one terminal node, don't you think?")
    valid <- FALSE
  }
  outs <- which((!is.root.node(ancestor) & !(ancestor %in% node)))
  if (length(outs) > 0) {
    str <- sprintf("the ancestor of node %s is not in the tree\n", node[outs])
    warning(str,call.=F)
    valid <- FALSE
  }
  anc <- ancestor.numbers(node,ancestor)
  ck <- all(
            sapply(
                   1:n,
                   function(x) {
                     good <- root %in% pedigree(anc,x)
                     if (!good) {
                       str <- sprintf("node %s is disconnected", node[x])
                       warning(str,call.=F)
                     }
                     good
                   }
                   )
            )
  valid && ck
}

# This file is part of the OUCH package.
# Author: Aaron A. King <king at tiem dot utk dot edu>
# It is distributed under the GNU Public License (see the file GPL
# included)
# The OUCH package is maintained at
#          http://www.tiem.utk.edu/~king/ouch/
#
