library(happy.hbrem)
library(multicore)
source("/data/www/preCC/R/happy.preCC.R")
fs.all <- function() {
  phenos = c("average.days.to.germ","days.to.germ","leaves.day.21","leaves.day.28","days.to.bolt","bolt.to.flower","basal","nodes","aerial","top.cm","total.cm","leaves.day.28.given.days.to.germ","basal.given.days.to.bolt","nodes.given.days.to.bolt","aerial.given.nodes")

  lapply(phenos, fs, mc.cores=8)
}

fs <- function(phenotype.name, phenotype.file="/data/scratch/rmott/ARABIDOPSIS/11092008/QTL_MAPPING_IGNORE/fitted.ignore.30112008.txt" , model="linear",chrs=1:19, logP.thresh=4, mc.cores=4,outfile=paste(phenotype.name, ".forward.selection.RData",sep=""), overwrite=FALSE, gscan.file=paste(phenotype.name,".gscandb.txt",sep=""), build="TAIR9", from.bp=NULL, to.bp=NULL, db=NULL ) {

#  if ( ! is.null(gscan.file) ) {
#    write.gscan.files(gscan.file, phenotype.name, build, chrs, logP.thresh)
#  }	

  if ( file.exists(outfile) & overwrite==FALSE) {
    cat(outfile, "exists\n")
    return(0);
  }
  pheno = read.delim( phenotype.file )
  if ( is.null(pheno[[phenotype.name]] ) ) {
    warning( "ERROR unknown phenotype ", phenotype.name, "\n")
    return(-1)
  }
  else {
    cat("analysing ", phenotype.name, "\n")
    y = pheno[[phenotype.name]]
    names(y) = pheno$CC.Line
    y = y[!is.na(y)]
    f.s = merge.analysis.forward.selection(chrs=chrs,y,phenotype.name,mc.cores=mc.cores,logP.thresh=logP.thresh, from.bp=from.bp, to.bp=to.bp, g=db)
    save(f.s,file=outfile)
		
    return(1)
  }
  
}

write.gscan.files <- function(gscan.file, phenotype.name, build, chrs=1:19, logP.thresh){
	df = NULL
	marker.data = NULL
	marker.mapping = NULL

	for( c in chrs ) {
		load(paste("chr",c, ".", phenotype.name, ".mergeanalysis.RData", sep=""))
		use = merge.analysis.result$logP.merge >logP.thresh
		merge.analysis.result = merge.analysis.result[use,]
		marker = paste( merge.analysis.result$chr, "-", merge.analysis.result$bp, sep="")
		merge.logP = merge.analysis.result$logP.merge 
	        marker.data =rbind(marker.data, data.frame(name=marker,marker_type="MERGE",leftseq="",rightseq="",alias=""))
		marker.mapping=rbind(marker.mapping,data.frame(marker=marker,genome_build=build,chromosome=merge.analysis.result$chr,bp_position=merge.analysis.result$bp,strand=1,cm=""))
		df = rbind(df, data.frame(marker=marker,additive.merge=merge.logP))
	}
	write.table(df,file=gscan.file,row=FALSE,quote=FALSE,sep="\t")
	write.table(marker.data,file="marker.csv",row=FALSE,quote=FALSE,sep=",")
	write.table(marker.mapping,file="marker.mapping.csv",row=FALSE,quote=FALSE,sep=",")
}

summarise.merge.analysis <- function( phenotype, chrs=1:19 ) {
	
	files = paste( "chr", chrs, ".", phenotype,  ".mergeanalysis.RData", sep="")
	results = NULL
	for( f in files ) {
		load(f)
		results = rbind(results, merge.analysis.result)
	}
	return(results)
}
merge.analysis.forward.selection <- function( chrs=1:19, y, label, model="linear", logP.thresh=4, mc.cores=8, from.bp=NULL, to.bp=NULL, g=NULL, condensed.dir="/data/www/preCC/preCC.merge/CONDENSED.X", merge.dir="/data/mouse.snps/" ) {

#  g = load.genome( "/data/scratch/rmott/ARABIDOPSIS/GENOME.CACHE.CLEAN/", chr=paste("chr",chrs,sep=""))

  if ( is.null(g) ) g = load.condensed.database( condensed.dir )
  use = intersect( g$subjects, names(y) )
  y.use = y[match(use, names(y), nomatch=0 )]

  d.use = match(use, g$subjects,  nomatch=0 )
  y.orig = y.use

  for( chr in chrs ) {
    chr = as.character(chr)
    cache.file = paste("chr", chr, ".", label, ".mergeanalysis.RData",sep="")
    if ( ! file.exists( cache.file ) ) {
      cat("merge.analysis.on.region", chr,"\n")
      load(paste(merge.dir, "chr", chr, ".merge.RData", sep=""))
      merge.analysis.result = merge.analysis.on.region( g, y, merge.data, markers=NULL, chr=chr, model=model, mc.cores=mc.cores, from.bp=from.bp, to.bp=to.bp)
      save(merge.analysis.result,file=cache.file)
      merge.analysis.result = NULL
      merge.data = NULL
      gc()
    }
  }

  genomewide.merge.data = list()
  genomewide.merge.results = list()
  n.candidates = 0
  best.variant = list(logP.merge=0)
  best.merge.df = NULL

  for( chr in chrs ) {
    chr = as.character(chr)
    cache.file = paste("chr", chr, ".", label, ".mergeanalysis.RData",sep="")
    load(paste(merge.dir, "chr", chr, ".merge.RData", sep=""))
    genomewide.merge.data[[chr]] = merge.data
    if ( file.exists( cache.file ) ) {
      cat("loading merge analysis cache", cache.file, "\n") 
      load(cache.file)
    }
    else {
      cat("merge.analysis.on.region", chr,"\n")
      merge.analysis.result = merge.analysis.on.region( g, y, genomewide.merge.data[[chr]], markers=NULL, chr=chr, model=model, mc.cores=mc.cores, from.bp=from.bp, to.bp=to.bp)
      save(merge.analysis.result,file=cache.file)
      merge.analysis.result = NULL
    }
    variants.to.use = sort(intersect( genomewide.merge.data[[chr]]$allele.data$bp, merge.analysis.result$bp))
    merge.analysis.result = merge.analysis.result[match(variants.to.use,merge.analysis.result$bp),]
    genomewide.merge.data[[chr]]$allele.data = genomewide.merge.data[[chr]]$allele.data[match(variants.to.use,genomewide.merge.data[[chr]]$allele.data$bp),]
    genomewide.merge.results[[chr]] = merge.analysis.result
    
    cat( chr, " merge.results " ,nrow(merge.analysis.result), "merge.data", nrow(genomewide.merge.data[[chr]]$allele.data), "\n")

    w = which.max(genomewide.merge.results[[chr]]$logP.merge)
    logP = genomewide.merge.results[[chr]]$logP.merge[w]
    if ( logP > best.variant$logP.merge ) {
      best.variant = genomewide.merge.results[[chr]][w,]
    }
    subset = genomewide.merge.results[[chr]]$logP.merge>logP.thresh
    if ( sum(subset) > 0 ) {
      genomewide.merge.data[[chr]]$allele.data = genomewide.merge.data[[chr]]$allele.data[subset,]
      n.candidates = n.candidates + nrow(genomewide.merge.data[[chr]]$allele.data)
    }
    else 
      genomewide.merge.data[[chr]]$allele.data = NULL
  }
  best.merge.df = rbind(best.merge.df, best.variant )
  print(best.merge.df)
  cat("candidates ", n.candidates, "\n")
  iter=1
  fitted.variants = NULL
  fitted.mat = list()
  mul=NULL
  while( n.candidates > 0 ) {

    best.result = merge.analysis.at.variant( y.use, best.variant, g, d.use, genomewide.merge.data, model=model)
    fitted.variants = rbind(fitted.variants$variant)
    fitted.mat[[best.result$id]] = best.result$mat
    mul = fit.multiple.model( y.orig, fitted.mat )  
    y.resid = resid(best.result$fit)
    best.df = best.result$fit$rank
    y.use = y.resid
    
    n.candidates = 0
    best.variant$logP.merge = 0
    for( chr in chrs ) {
      if ( ! is.null(genomewide.merge.data[[chr]]$allele.data)) {
        genomewide.merge.results[[chr]] = merge.analysis.on.region( g, y.use, genomewide.merge.data[[chr]], markers=NULL, chr=chr, model=model, mc.cores=mc.cores, from.bp=from.bp, to.bp=to.bp)
        cat( chr, " merge.results " ,nrow(genomewide.merge.results[[chr]]), "merge.data", nrow(genomewide.merge.data[[chr]]$allele.data), "\n")
        print(class(genomewide.merge.results[[chr]]))
        if ( nrow(genomewide.merge.results[[chr]]) != nrow(genomewide.merge.data[[chr]]$allele.data) ) 

        w = which.max(genomewide.merge.results[[chr]]$logP.merge)
        logP = genomewide.merge.results[[chr]]$logP.merge[w]
        if ( logP > best.variant$logP.merge ) {
          best.variant = genomewide.merge.results[[chr]][w,]
        }
        subset = genomewide.merge.results[[chr]]$logP.merge>logP.thresh
        if ( sum(subset) > 0 ) {
          genomewide.merge.data[[chr]]$allele.data = genomewide.merge.data[[chr]]$allele.data[subset,]
          n.candidates = n.candidates + nrow(genomewide.merge.data[[chr]]$allele.data)
        }
      }
      iter = iter+1
    }
    best.merge.df = rbind(best.merge.df, best.variant )
    print(best.merge.df)
    cat("candidates ", n.candidates, "\n")
  }
  return( list( forward.selection=best.merge.df, multiple.model.fit=mul, fitted.mat=fitted.mat) )
}

fit.multiple.model <- function( y, fitted.mat )  {
  fit.formula = as.formula(paste("y ~ ", paste(  names(fitted.mat), collapse=" + ")))
  f = lm(fit.formula,data=fitted.mat,y=TRUE)
  a = anova(f)
  print(a)
  return(f)
}
    
  

make.merge.factors <- function(chrs=c(1:19,"X"), suffix=".CC.txt", mc.cores=8) {
  mclapply( chrs, function( c, suffix ) {
    f = paste("chr", c, suffix, sep="")
    cat("reading ", f, "\n")
    variants = read.table(f, h=T, sep="\t")
    merge.data=merge.factors(variants)
    save(merge.data,file=paste("chr", c, ".merge.RData", sep=""))
  }, suffix, mc.cores=mc.cores )
}
    
merge.factors<- function( allele.data ) {
  alleles = t(apply(allele.data[,3:ncol(allele.data)],1,as.character))
  alleles.integer = t(apply( alleles , 1,
        function( X ) {
          return(as.integer(factor( X, levels=unique(X) )))
        }
    ))
  
  sdp = factor(apply( alleles.integer,1, function( X ) { paste( X, collapse=".") }))
  sdp.unique = sort(unique(sdp))
  sdp.index = match(sdp.unique, sdp)
  matrices = apply( alleles.integer[sdp.index,], 1, function(X) { mat = sapply (unique(X), function( y, X) { ifelse(y==X, 1, 0) }, X) } )
  names(matrices) = sdp.unique
  allele.data $sdp = sdp
  return( list( allele.data=allele.data, matrices=matrices))
}

# allele.data is table with columns
#chr	bp	A/J	C57BL/6J	129S1SvImJ	NOD/LtJ	NZO/HILtJ	CAST/EiJ	PWK/PhJ	WSB/EiJ



  
log0 <- function(x) { ifelse( x>0, log(x), 0 )}

lm.merge.analysis.on.sdp <- function( sdp, d, mat, y, fit.null, fit.d, poe.data=NULL ) {
  m = mat[[sdp]];
  if ( !is.null(poe.data) ) {
	browser()
	m.poe = cbind(m,m)
	d.additive.poe = poe.data$additive.poe %*% m.poe
	d.additive = poe.data$additive %*% m
	fit.merge.poe = lm( y ~ d.additive.poe )
	fit.merge.additive = lm( y ~ d.additive )
	a = anova( fit.null, fit.merge.additive, fit.merge.poe, fit.d )
	a1 = anova(fit.null, fit.merge.additive)
	a2 = anova(fit.null, fit.merge.poe )
	return( c(a[4,6], a[3,6], a[2,6], a1[2,6], a2[2,6])) 
# partial poe vs merge.poe
# partial merge.poe vs unmerged.additve.poe
# unmerged
  }
  else {
	  dd = d %*% m;
	  ent = mean(apply(dd,1,function(p) { -sum(p*log0(p))}))
    
	  fit.dd = lm( y ~ dd )
	  a = anova( fit.null, fit.dd, fit.d)
#  cat("sdp", sdp, -log10(a[3,6]), -log10(a[2,6]), "\n")
	  return(c(a[3,6],a[2,6],ent))
  }
}

glm.merge.analysis.on.sdp <- function( sdp, d, mat, y, fit.null, fit.d, poe.mat=NULL ) {
  m = mat[[sdp]];
  if ( !is.null(poe.mat) ) 
	m = cbind(m,m)
  dd = d %*% m;
  ent = mean(apply(dd,1,function(p) { -sum(p*log0(p))}))
  fit.dd = glm( y ~ dd, family="binomial" )
  a = anova( fit.null, fit.dd, fit.d, test="Chisq")
#  cat("sdp", sdp, -log10(a[3,6]), -log10(a[2,6]), "\n")
  return(c(a[3,5],a[2,5],ent))
}

merge.analysis.at.variant <- function( y, variant, g, d.use, genomewide.merge.data, model="linear" ) {
  snp = as.character(variant$snp)
  chr = as.character(variant$chr)
  bp = variant$bp
  sdp = variant$sdp
  m = genomewide.merge.data[[chr]]$matrices[[sdp]]
#  d = hdesign(g,snp)

  d = condensed.hdesign(g,snp)
  d = d[d.use,]
  
  if ( model=="linear") {
    fit.null = lm( y ~ 1 )
    fit.d = lm( y ~ d )
    
    dd = d %*% m;
    fit.dd = lm( y ~ dd )
    a = anova(fit.null, fit.dd,fit.d)
    return( list(fit=fit.dd, anova=a, variant=variant, id=paste("variant.", chr,".", bp,sep=""), mat=dd))
  }
  else if ( model == "binary" ) {
    warning("not implemented")
    fit.null = glm( y ~ 1, family="binomial" )
    fit.d = glm( y ~ d, family="binomial" )
    a = anova(fit.null, fit.d, test="Chisq")
    lpval = -log10(a[2,5])
  }
}

  
merge.analysis.on.locus <- function( y, d, snp, chr, from, to, merge.data, model="linear", poe.data=NULL ) {

  lpval = NA
  if ( model=="linear") {
    fit.null = lm( y ~ 1 )
    fit.d = lm( y ~ d )
    a = anova(fit.null, fit.d)
    if ( is.numeric(a[2,6])) lpval = -log10(a[2,6])
    else browser()
  }
  else if ( model == "binary" ) {
    fit.null = glm( y ~ 1, family="binomial" )
    fit.d = glm( y ~ d, family="binomial" )
    a = anova(fit.null, fit.d, test="Chisq")
    if ( is.numeric(a[2,5])) lpval = -log10(a[2,5])
  }
    
  allele.data = merge.data$allele.data
  variants = allele.data[allele.data$chr==chr & allele.data$bp >= from & allele.data$bp <=to,]

  nv = nrow(variants)
  logP = NULL
  if ( !is.null(nv) & nv>0 ) {
    variants$sdp = as.factor(as.character(variants$sdp))
    sdp.unique = as.character(levels(variants$sdp))

    zero = rep(0.0,nv)
    if ( is.null(poe.mat)
	    logP = data.frame(chr=rep(chr,nv), bp=variants$bp, bp=variants$bp, snp=rep(snp, nv), "logP.interval" = rep(lpval,nv), "logP.merge"=zero, "logP.partial"=zero,"sdp"=rep("",nv))
    else 
	    logP = data.frame(chr=rep(chr,nv), bp=variants$bp, bp=variants$bp, snp=rep(snp, nv), "logP.interval" = rep(lpval,nv), "logP.merge"=zero, "logP.partial"=zero,logP.interval.poe = zero, logP.merge.poe=zero, logP.merge.poe.partial."sdp"=rep("",nv))

    if ( model=="linear") {
      t.tmp = t(sapply(sdp.unique, lm.merge.analysis.on.sdp, d, merge.data$matrices, y,  fit.null, fit.d, poe.data=poe.data ))
      if ( sum(!is.numeric(t.tmp)) > 0 )
        browser()
      logP.unique = t.tmp
    }
    else {
      logP.unique = -t(sapply(
                                    sdp.unique, 
                                    glm.merge.analysis.on.sdp, d, merge.data$matrices, y,  fit.null, fit.d, poe.data=poe.data ))
    }
    idx = as.integer(variants$sdp)
    if ( is.null(poe.data)) {   
	
	    logP$logP.partial = -log10(logP.unique[idx,1])
	    logP$logP.merge = -log10(logP.unique[idx,2])
	    logP$entropy = logP.unique[idx,3]
	    logP$sdp = as.character(variants$sdp)
	}	
    else {
	logP$logP.partial 
  }
  return( logP )
}


merge.analysis.on.region <- function( g, y, merge.data, markers=NULL, chr=NULL, from.bp=NULL, to.bp=NULL, model="additive", poe=FALSE, mc.cores=8 ) {

  poe.mat = NULL
  if ( poe == TRUE ) {
	if ( class(g) == "happy" ) 
		poe.mat = additive.matrix(g$strains)
	else {
		error( "poe analysis not possible\n")
		return(NULL)
	}
  }

  use = intersect( g$subjects, names(y) )
  y.use = y[match(use, names(y), nomatch=0 )]
  d.use = match(use, g$subjects,  nomatch=0 )

  if ( is.null(from.bp)) 
	from.bp = 0
  if ( is.null(to.bp))
	to.bp = 1.0e20;
	  
  if ( is.null(markers) ){
    if ( is.null(chr) )
      markers = as.character(g$additive$genome$marker)
    else {
      markers = as.character(g$additive$genome$marker[g$additive$genome$chr==chr & g$additive$genome$bp >= from.bp & g$additive$genome$bp <= to.bp ])

      m = match(markers[1], g$additive$genome$marker)
      if ( m > 1 ) 
	if (g$additive$genome$chr[m-1] == chr ) markers = c( as.character(g$additive$genome$marker[m-1]), markers)
    }
  }
  else if ( is.numeric(markers[1]))
    markers = as.character(g$additive$genome$marker[markers])
  m.index = match(markers, g$additive$genome$marker)
  if ( mc.cores > 1 ) 
    merge.analysis = mclapply( m.index, merge.analysis.locus, g,  y.use, d.use, merge.data, model, poe.mat, mc.cores=mc.cores )
  else
    merge.analysis = lapply( m.index, merge.analysis.locus, g,  y.use, d.use, merge.data, model, poe.mat )

  res = do.call("rbind", merge.analysis)
  if ( class(res) != "data.frame")
    browser()
    res = res[match(unique(res$bp),res$bp),]
  return(res)
}

merge.analysis.locus <-     function( m, g, y.use, d.use, merge.data, model, poe.mat ) {
#  cat(m, " marker ", g$additive$genome$marker[m], "\n");
#  d = hdesign(g,m)

  if ( is.null(poe.mat) ) {
	 d = condensed.hdesign(g,as.character(g$additive$genome$marker[m]))
	 d = d[d.use,]
	 poe.data = NULL
  }
  else {
	d = hprob2(g,as.character(g$additive$genome$marker[m]))
	d.additive = d %*% poe.mat$additive
	d.additive = d.additive[d.use,]
	d.poe.additive = d %*% poe.mat$poe.additive
	d.poe.additive = d.poe.additive[d.use,]
	poe.data = list(d.poe.additive=d.poe.additive, d.additive=d.additive)
  }	
  info = g$additive$genome[m,]
  snp = info$marker;	
  chr = info$chr	
  from = info$bp
  to = -1
  if ( m < nrow(g$additive$genome) ) {
  to = g$additive$genome$bp[m+1]
  }
  if ( to < from ) 
    to = from + 1.0e7
  return( merge.analysis.on.locus( y.use, d, snp,  chr, from, to, merge.data, model=model, poe.data=poe.data )) 
}

merge.analysis.all<- function( g, phenotype.name, phenotype.data,model="linear", from.bp=NULL, to.bp=NULL, merge.dir="/data/mouse.snps/" ) {

  for( chr in 1:19 ) {
    y = phenotype.data[,phenotype.name]
    names(y) = phenotype.data$CC.Line
    load(paste(merge.dir, "chr", chr, ".merge.RData", sep=""))
#    merge.data = get(paste("chr", chr, ".merge", sep=""))
    assign( paste("chr", chr, ".merge", sep=""), merge.data )
    cat( "merging ", phenotype.name, " ", chr, "\n")
    merge.results = merge.analysis.on.region( g, y, merge.data, markers=NULL, chr=chr, model=model, from.bp=from.bp, to.bp=to.bp)
    write.table(merge.results, file=paste( phenotype.name, ".chr", chr, ".merge.txt", sep=""), quote=FALSE, row=FALSE )
  }
}

merge.genome <- function(phenotype.name, file="/data/scratch/rmott/ARABIDOPSIS/11092008/QTL_MAPPING_IGNORE/fitted.ignore.30112008.txt" , model="linear" ) {
  pheno = read.delim( file )
  if ( is.null(pheno[[phenotype.name]] ) ) {
    warning( "ERROR unknown phenotype ", phenotype.name, "\n")
  }
  else {
    g = load.genome( "/data/scratch/rmott/ARABIDOPSIS/GENOME.CACHE.CLEAN/", chr=paste("chr",1:19,sep=""))
    merge.analysis.all( g, phenotype.name, pheno, model=model )
  }
}

additive.matrix <- function( strains ) {
   nstrains = length(strains)
   add.mat1 = matrix( 0, nrow=nstrains*nstrains, ncol=nstrains )
   add.mat2 = matrix( 0, nrow=nstrains*nstrains, ncol=nstrains )
   for(i in 1:nstrains) {
     j = 1+nstrains*(i-1)
     add.mat1[j:(j+nstrains-1),i] = 1
     k = seq(i,i+nstrains*(nstrains-1), nstrains)
     add.mat2[k,i] = 1
   }
   poe.additive = cbind(add.mat1,add.mat2)/2
   colnames(poe.additive) = c(paste("maternal.", strains,sep=""), paste("paternal.", strains,sep=""))
   additive = (add.mat1 + add.mat2)/2
   colnames(additive) = strains
   return(list(additive=additive,poe.additive=poe.additive))
}
  

