BAliPhy.Main.hs

{-# LANGUAGE ExtendedDefaultRules #-}
{-# LANGUAGE OverloadedStrings #-}

module Main where

import BAliPhy.Run
import Bio.Alignment
import Bio.Alphabet
import Bio.Sequence
import qualified Data.IntMap as IntMap
import Data.JSON ((.=))
import qualified Data.JSON as J
import qualified Data.Map as Map
import qualified Data.Set as Set
import qualified Data.Text.IO as T
import IModel
import MCMC
import Options.Applicative
import Probability
import Probability.Logger
import Probability.Random
import SModel
import SModel.Parsimony
import System.Exit
import System.FilePath
import System.IO
import Tree
import Tree.Newick

sample_smodel alpha = do
    sym <- sample (symmetricDirichletOn (Set.fromList (letter_pair_names alpha)) 1)
    pi <- sample (symmetricDirichletOn (letterSet alpha) 1)
    alpha_2 <- sample (logLaplace 6 2)
    pInv <- sample (uniform 0 1)
    r <- sample (gamma 2 (1 / 4))
    pi1 <- sample (beta 2 2)
    let rate = ((2 * r) * pi1) * (1 - pi1)
    let s01 = r * pi1
    let s10 = r * (1 - pi1)
    let result =
            ((((gtr' sym pi alpha) +> unitMixture) +> (SModel.gammaRatesOn alpha_2 4)) +> (plusInv pInv)) +> (huelsenbeck02 s01 s10)
    let loggers =
            LoggerValues
                [ "GTR:sym" %=% sym
                , "GTR:pi" %=% pi
                , "ASRV.Gamma:alpha" %=% alpha_2
                , "Inv:pInv" %=% pInv
                , "Covarion.Huelsenbeck02:r" %=% r
                , "Covarion.Huelsenbeck02:pi1" %=% pi1
                , "Covarion.Huelsenbeck02:rate" %=% rate
                , "Covarion.Huelsenbeck02:s01" %=% s01
                , "Covarion.Huelsenbeck02:s10" %=% s10
                ]
                (contextFields [])
    return (result, loggers)

sample_imodel topology = do
    rate <- sample (logLaplace (negate 4) 0.70699999999999996)
    meanLength <- sample (shifted_exponential 10 1)
    let result = IModel.rs07 rate meanLength topology
    let loggers = LoggerValues ["RS07:rate" %=% rate, "RS07:meanLength" %=% meanLength] (contextFields [])
    return (result, loggers)

sample_scale = sample (gamma 0.5 2)

sampleTree taxa = sample (uniformLabelledTree taxa (gamma 0.5 (1 / (intToDouble (length taxa)))))

model sequenceData logParamsTSV logParamsJSON logTree [logA] [logCatStates] = do
    let taxa = getTaxa sequenceData
    tree <- sampleTree taxa
    (indelRates, log_indelRates) <- do
        sigma <- sample (logLaplace (negate 3) 1)
        xs <- sample (iidMap (getUEdgesSet tree) (logNormal 0 1))
        let result = fmap (\x_4 -> x_4 Prelude.** sigma) xs
        let loggers = LoggerValues ["sigma" %=% sigma] (contextFields [])
        return (result, loggers)
    let indelTree = addBranchRates indelRates tree
    let tlength = treeLength tree
    scale1 <- sample_scale
    addMove 2 (scaleGroupsSlice [scale1] (branchLengths tree))
    addMove 1 (scaleGroupsMH [scale1] (branchLengths tree))
    (smodel, log_smodel) <- sample_smodel rna
    (imodel, log_imodel) <- sample_imodel tree
    let sequence_lengths = getSequenceLengths sequenceData
    (alignment, properties_A) <- sampleWithProps (phyloAlignment indelTree imodel scale1 sequence_lengths)
    properties <- observe sequenceData (phyloCTMC tree alignment smodel scale1)
    let alignment_length = alignmentLength alignment
    let num_indels = totalNumIndels alignment
    let total_length_indels = totalLengthIndels alignment
    let prior_A = ln (probability properties_A)
    let ancStates = prop_anc_cat_states properties
    let ancAlignment = toFasta (ancestralAlignment tree alignment (getSMap smodel) rna ancStates)
    let catStates = labeledNodeMap tree ancStates
    let substs = parsimony tree (unitCostMatrix rna) (sequenceData, alignment)
    let part1Loggers =
            [ "|A|" %=% alignment_length
            , "#indels" %=% num_indels
            , "|indels|" %=% total_length_indels
            , "prior_A" %=% prior_A
            , "likelihood" %=% (ln (prop_likelihood properties))
            , "#substs" %=% substs
            ]
    let alignmentLengths = [alignment_length]
    let scale = scale1
    let loggerValues =
            LoggerValues
                [ "indelRates" %>% (parameterLogValues log_indelRates)
                , "|T|" %=% tlength
                , "scale1" %=% scale1
                , "scale1*|T|" %=% (scale1 * tlength)
                , "S1" %>% (parameterLogValues log_smodel)
                , "I1" %>% (parameterLogValues log_imodel)
                , "P1" %>% part1Loggers
                , "scale" %=% scale
                , "scale*|T|" %=% (scale * tlength)
                , "|A|" %=% (sum alignmentLengths)
                , "#indels" %=% (sum [num_indels])
                , "|indels|" %=% (sum [total_length_indels])
                , "#substs" %=% (sum [substs])
                , "prior_A" %=% (sum [prior_A])
                ]
                (contextFields (["prior" %=! logPrior, "likelihood" %=! logLikelihood, "posterior" %=! logPosterior] ++ []))
    addLogger $ (logParamsTSV loggerValues)
    addLogger $ (logParamsJSON loggerValues)
    addLogger $ (logTree (addInternalLabels (scaleBranchLengths scale tree)))
    addLogger $ ((every 10) $ (logA ancAlignment))
    addLogger $
        ( (every 10) $
            ( logCatStates
                ( (((J.toJSONKey "catStates") .= catStates) <> ((J.toJSONKey "properties") .= (prop_smodel_properties properties)))
                    <> ((J.toJSONKey "conditions") .= (prop_smodel_conditions properties))
                )
            )
        )
    return (parameterLogValues loggerValues)

runOptions =
    info
        (modelRunOptions "25" 200000 [TSV] <**> helper)
        (fullDesc <> progDesc "Run this generated BAli-Phy analysis")

-- Test mode never evaluates logger paths, so an empty placeholder keeps logger setup uniform.
modelRunDirectory TestRun = ""
modelRunDirectory (MCMCRun directory) = directory

-- Describe a logger only after its destination has been opened successfully.
reportOutput description filename suffix =
    putStrLn $ "   - Sampled " ++ description ++ " logged to " ++ show filename ++ suffix

main = do
    options <- execParser runOptions
    runInfo <- initializeModelRun (runMode options)
    let isTest = (runMode options) == TestMode
    let loggingEnabled = not isTest
    let tsvEnabled = loggingEnabled && (elem TSV (logFormats options))
    let jsonEnabled = loggingEnabled && (elem JSON (logFormats options))
    let outputDirectory = modelRunDirectory runInfo
    sequenceData <- (mkUnalignedCharacterData rna) <$> (loadSequences "25.fasta")
    logParamsTSV <- if tsvEnabled then tsvLogger (outputDirectory </> "C1.log") ["iter"] else return noLogger
    logParamsJSON <- if jsonEnabled then jsonLogger (outputDirectory </> "C1.log.json") else return noLogger
    logTree <- if loggingEnabled then treeLogger (outputDirectory </> "C1.trees") else return noLogger
    logA <- if loggingEnabled then alignmentLogger (outputDirectory </> "C1.P1.fastas") else return noLogger
    logCatStates <-
        if loggingEnabled then ejsonLogger (outputDirectory </> "C1.P1.site-property-samples.jsonl") else return noLogger
    unless
        isTest
        ( do
            putStrLn ""
            putStrLn "Beginning MCMC computations."
            when tsvEnabled (reportOutput "numerical parameters" (outputDirectory </> "C1.log") " as TSV")
            when jsonEnabled (reportOutput "numerical parameters" (outputDirectory </> "C1.log.json") " as JSON")
            when loggingEnabled (reportOutput "trees" (outputDirectory </> "C1.trees") "")
            when loggingEnabled (reportOutput "alignments" (outputDirectory </> "C1.P1.fastas") "")
            when loggingEnabled (reportOutput "character properties" (outputDirectory </> "C1.P1.site-property-samples.jsonl") "")
            putStrLn ""
            putStrLn "BAli-Phy does NOT detect how many iterations is sufficient:"
            putStrLn "   You need to monitor convergence and kill it when done."
            putStrLn ("   Maximum number of iterations set to " ++ ((show (iterations options)) ++ "."))
            putStrLn ""
            when
                tsvEnabled
                ( putStrLn
                    "You can examine 'C1.log' using BAli-Phy tool statreport (command-line) or the BEAST program Tracer (graphical)."
                )
            putStrLn "See the manual at http://www.bali-phy.org/README.html for further information."
            hFlush stdout
        )
    mcmcState <- makeMCMCState $ (model sequenceData logParamsTSV logParamsJSON logTree [logA] [logCatStates])
    if isTest then printInitialModel (logFormats options) mcmcState else runMCMC (iterations options) mcmcState