{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "faf49fef-55a7-4932-ac97-e50c8bd9f504",
   "metadata": {},
   "source": [
    "## answers 13: the adventure of the three trees"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "8b87bc29-e43f-42ef-9888-5896c9e6b5e4",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "43a343bb-7ad7-4a60-b1a8-2e0dc4098939",
   "metadata": {},
   "source": [
    "### 1. implement code for simulating sequences down a tree"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "d2ae627f-3ee8-4527-943b-dcdbd569eb7c",
   "metadata": {},
   "source": [
    "First, let's instantiate the three pset trees in a Python data structure.\n",
    "\n",
    "I'll write a function `specify_trees()` that returns the three trees as a list. Tree T0 is `T[0]`, and so on. \n",
    "\n",
    "Each tree is a list of 2n-1 nodes numbered 0..2n-2 for n taxa. Nodes 0..n-1 are the n observed leaf taxa. Nodes n..2n-2 are the internal nodes. Node 2n-2 is the root. \n",
    "\n",
    "Each node is set to either `None` (for leaf taxa that have no children) or a tuple of (left_child, right_child). Then each child is itself a tuple of (child_idx, branch_len) to identify which node is the child, and the length of the branch to it in substitutions/site."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "d688cdd3-8c76-49de-9769-12b25310570b",
   "metadata": {},
   "outputs": [],
   "source": [
    "def specify_trees():\n",
    "    r\"\"\"\n",
    "    Create the three trees for the problem (\"pset trees\").\n",
    "\n",
    "    Args:\n",
    "      (none)\n",
    "\n",
    "    Returns:\n",
    "      List of three tree structures.\n",
    "\n",
    "    Trees[0]  = true tree... ML, NJ will favor\n",
    "    Trees[1]  = long branch attraction tree... UPGMA, parsimony will favor\n",
    "    Trees[2]  = the remaining topology...  no method will favor\n",
    "\n",
    "         .6            .6             .6\n",
    "        / \\           / \\            / \\\n",
    "       .4  .5        .4  .5         .4  .5\n",
    "      /|   |\\       /|   |\\        /|   |\\\n",
    "     / 0   1 \\     0 1   | \\      / 0   1 \\\n",
    "    2         3          2  3    3         2\n",
    "  \n",
    "     T0 = true       T1 = LBA      T2 = other\n",
    "\n",
    "    A (rooted) tree is a list of 2N-1 nodes, 0..2N-2. \n",
    "       Nodes 0..N-1 are the taxa on the leaves.\n",
    "       These leaf nodes are set to None: they have no children.\n",
    "       Nodes N..2N-2 are the internal nodes. 2N-2 is the root.\n",
    "    Each internal node is a tuple of (left_child, right_child).\n",
    "    Each child is a tuple of (idx, branchlen).\n",
    "       idx = 0..2N-2 index of child node.\n",
    "       branchlen = branch length to that child, in substitutions/site\n",
    "        \n",
    "    A branch length in units of substitutions/site is 3 \\alpha t in Jukes-Cantor:\n",
    "    so \\alpha t = branchlen/3, when we go to parameterize substitution probs from\n",
    "    these branch lengths.\n",
    "    \"\"\"\n",
    "\n",
    "    #         --- leaf taxa (0..3) ----  - internal node (4) -  - internal node (5) -  --- root node (6) ---\n",
    "    Trees = [ [ None, None, None, None,  ((2, 0.6), (0, 0.1)),  ((1, 0.1), (3, 0.6)),  ((4, 0.05), (5, 0.05))],\n",
    "              [ None, None, None, None,  ((0, 0.1), (1, 0.1)),  ((2, 0.6), (3, 0.6)),  ((4, 0.05), (5, 0.05))],\n",
    "              [ None, None, None, None,  ((3, 0.6), (0, 0.1)),  ((1, 0.1), (2, 0.6)),  ((4, 0.05), (5, 0.05))] ]\n",
    "    return Trees"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "29ef5530-44f5-41cb-889c-0971fc835bcc",
   "metadata": {},
   "source": [
    "Now I want a function `sample_sequences()` to generate sequences down a tree.\n",
    "\n",
    "I'll store my multiple sequence alignment as a nseq x L ndarray, with DNA residues digitized as values 0..3 representing A..T.\n",
    "\n",
    "I'll write a function `sample_descendant()` to generate one child sequence from one ancestor, and  a function `p_jukescantor()` to give me the $P(b \\mid a,t)$ substitution probabilities as a matrix `P[a][b]`, given a branch length.\n",
    "\n",
    "A potentially tricky bit here is that in a Jukes-Cantor model, the expected number of substitutions/site is $3 \\alpha t$. So if we're given a branch length $b_i$ in units of substitutions/site, $\\alpha t = \\frac{b_i}{3}$. "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "954c2727-3f04-49fc-b854-9e1d3c5a6460",
   "metadata": {},
   "outputs": [],
   "source": [
    "def p_jukescantor(branchlen):\n",
    "    r\"\"\"Create Jukes/Cantor P(b|a,t) given branch len in subst/site.\n",
    "\n",
    "    Args:\n",
    "        branchlen (float >= 0) : branch length in subst/site\n",
    "\n",
    "    Returns:\n",
    "        P (ndarray, 4x4) : P[a,b] is P(b|a,t), where t is branchlen/3\n",
    "    \"\"\"\n",
    "    rp = 1/4 + 3/4 * np.exp(-4/3 * branchlen)\n",
    "    sp = 1/4 - 1/4 * np.exp(-4/3 * branchlen)\n",
    "    P = np.array([[ rp, sp, sp, sp ],\n",
    "                  [ sp, rp, sp, sp ],\n",
    "                  [ sp, sp, rp, sp ],\n",
    "                  [ sp, sp, sp, rp ]])\n",
    "    return P"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "f8428a0c-8881-4899-ab0a-ba55f96b7953",
   "metadata": {},
   "outputs": [],
   "source": [
    "def sample_descendant(rng, Y, P):\n",
    "    r\"\"\"Sample a descendant sequence X, given ancestor Y.\n",
    "\n",
    "    Args:\n",
    "        rng              : numpy RNG\n",
    "        Y (ndarray, L)   : ancestral sequence\n",
    "        P (ndarray, 4x4) : P[a,b] is P(b|a,t); t is blen/3 for Jukes/Cantor\n",
    "\n",
    "    Returns:\n",
    "        X (ndarray, L)   : descendant sequence\n",
    "    \"\"\"\n",
    "    L = len(Y)       # ancestral sequence length\n",
    "    K = len(P[0])    # alphabet size\n",
    "\n",
    "    X = np.zeros( L, dtype=int)\n",
    "    for i in range(L):\n",
    "        X[i] = rng.choice(K, p=P[Y[i]])\n",
    "    return X\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "id": "e9b98fda-3bd3-4a75-b836-b396fc7c9c05",
   "metadata": {},
   "outputs": [],
   "source": [
    "def sample_sequences(rng, T, L):\n",
    "    r\"\"\"\n",
    "    Sample a multiple seq alignment down a tree.\n",
    "\n",
    "    Args:\n",
    "        rng  : numpy RNG\n",
    "        T    : rooted tree structure to generate down\n",
    "        L    : length of sequences to generate\n",
    "\n",
    "    Returns:\n",
    "        msa  : (nnodes x L) ndarray of digital seqs.\n",
    "               Values are 0..3 for digitized DNA residues A..T.\n",
    "               Rows 0..ntaxa-1 are seqs at the leaves.\n",
    "               Rows ntaxa..nnodes-1 are ancestral seqs at internal nodes\n",
    "               Row nnodes-1 is the last common ancestor, at root.\n",
    "\n",
    "    General; works on any tree T, not just a pset tree.\n",
    "    \"\"\"\n",
    "    nnodes = len(T)\n",
    "    ntaxa  = (nnodes + 1)//2\n",
    "    rooty  = nnodes-1\n",
    "\n",
    "    msa = np.zeros( (nnodes, L), dtype=int)   # msa[0..n-1] = leaves; msa[n..2n-2] = internal; msa[2n-2] = root\n",
    "\n",
    "    msa[rooty] = rng.integers(0, 4, L)        # generate a random ancestral digital DNA sequence at the root\n",
    "    stack    = [ rooty ]                      # initialize a stack with the root idx on it\n",
    "    while len(stack) > 0:                     # iterate through internal nodes, root to leaves:\n",
    "        a      = stack.pop()                  # a = idx of ancestor node. We know a is internal because we only push internal nodes\n",
    "        for (b, blen) in T[a]:                \n",
    "            P    = p_jukescantor(blen)        # branch len is in subst/site: 3at in jukes-cantor time t units\n",
    "            msa[b] = sample_descendant(rng, msa[a], P)\n",
    "            if b >= ntaxa: stack.append(b) \n",
    "    return msa"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "25fb65ab-84ed-4440-91f2-42433c41a522",
   "metadata": {},
   "source": [
    "### 2. implement mini UPGMA"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "be5faff1-812f-42d5-8b1c-11fbb5ae2245",
   "metadata": {},
   "source": [
    "For the distance-based methods, we need to first calculate a pairwise distance matrix $d_{ij}$ for all pairs of sequences $(i,j)$. \n",
    "\n",
    "I'll do that in two pieces. First, a function to calculate the Jukes-Cantor distance for one pair of sequences; second, a function to run that over all pairs $(i,j)$.\n",
    "\n",
    "A potentially tricky part here is what happens if the observed fractional difference (the so-called \"p distance\") between two sequences is $\\geq$ 0.75. This can happen by stochastic chance for finite-length DNA sequences that are very distantly related; for example for taxa (2,3) here. The Jukes-Cantor distance is then undefined. We can catch this case and assign an infinite distance, but then we'll create other problems with subtracting infinities and getting NaN's. It suffices as a hack here to just assign a ridiculously large but finite distance: here I'm assigning 100."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "id": "7d66915f-ba14-4509-a219-bc5260e257f9",
   "metadata": {},
   "outputs": [],
   "source": [
    "def dist_jukescantor(dsq1, dsq2):\n",
    "    r\"\"\"Calculate Jukes/Cantor distance between two aligned seqs.\n",
    "\n",
    "    Args:\n",
    "        dsq1 (ndarray, L) : digital sequence 1, from ungapped MSA\n",
    "        dsq2 (ndarray, L) : digital sequence 2\n",
    "\n",
    "    Returns:\n",
    "        (float, >= 0) : estimated J/C distance in subst/site\n",
    "\n",
    "    Let f = fractional difference = nsubst / (nsubst + nid).\n",
    "    Then d = -3/4 log(1 - 4/3 f).\n",
    "\n",
    "    When f >= 0.75, Jukes/Cantor distance is undefined, but f >= 0.75\n",
    "    can easily happen for distantly related sequences. To handle\n",
    "    this case, we set an arbitrarily large distance of 100. Though\n",
    "    arbitrary, this is better than setting an element to infinity,\n",
    "    because subsequent operations on infinite distances can result\n",
    "    in NaN's and undefined bad behavior.\n",
    "\n",
    "    d is a branch length in subst/site.\n",
    "    \"\"\"\n",
    "    nid, nsubst = 0, 0\n",
    "    alen        = len(dsq1)\n",
    "    assert(len(dsq2) == alen)\n",
    "    \n",
    "    for a,b in zip(dsq1, dsq2):\n",
    "        if a == b: nid    += 1\n",
    "        else:      nsubst += 1\n",
    "\n",
    "    fsubst = nsubst / (nid + nsubst)\n",
    "    if fsubst < 0.75: d = -3/4 * np.log(1 - fsubst*4/3)\n",
    "    else:             d = 100.          # arbitrary large distance\n",
    "    return d "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "id": "11fb9605-3fb0-4a72-bf68-00f7cd55b669",
   "metadata": {},
   "outputs": [],
   "source": [
    "def dist_matrix(msa):\n",
    "    r\"\"\"Calculate Jukes/Cantor distance matrix d_ij for all seq pairs in msa.\n",
    "\n",
    "    Args:\n",
    "        msa (ndarray, nseq x L) : digitized ungapped mult seq alignment\n",
    "\n",
    "    Returns:\n",
    "       (ndarray, nseq x nseq) : d_ij for all pairs of seqs\n",
    "\n",
    "    The msa we generate for the pset contains both leaf taxa and\n",
    "    internal ancestors. We usually only need to calculate d_ij for\n",
    "    the taxa. Therefore this function is usually called as something\n",
    "    like `dist_matrix(msa[:ntaxa])`.\n",
    "\n",
    "    d_ij is symmetric (d[i,j] == d[j,i]) with 0's down diagonal (d[i,i] = 0).\n",
    "    \"\"\"\n",
    "    nseq   = len(msa)\n",
    "    ncol   = len(msa[0])\n",
    "\n",
    "    d = np.zeros( (nseq, nseq) )   # diagonal i=j initialized to zero here. Others replaced below.\n",
    "    for i in range(nseq):\n",
    "        for j in range(i+1, nseq):\n",
    "            d[i,j] = d[j,i] = dist_jukescantor(msa[i], msa[j])\n",
    "    return d"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "accb94a0-02f8-478e-a602-648bfcbf626d",
   "metadata": {},
   "source": [
    "Now I'll write a function `best_distance_tree()` that I'll share between mini UPGMA and mini NJ, where I can use the same function to finding argmin (i,j) in the $d_{ij}$ matrix directly (for UPGMA) or find argmin (i,j) in the NJ modified $D_{ij}$ matrix). Then, given the best (i,j), I know which unrooted tree topology is supported.\n",
    "\n",
    "A potentially tricky thing here is, what to do about ties, when more than one (i,j) pair has the same minimum distance. Here I go to the trouble of randomly choosing one of the optimal solutions."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "id": "20c54718-c284-42c9-93b2-9a6d5c4c9c93",
   "metadata": {},
   "outputs": [],
   "source": [
    "def best_distance_tree(rng, D):\n",
    "    r\"\"\"Find which pset tree is supported by distance matrix D.\n",
    "\n",
    "    Args:\n",
    "        rng : NumPy RNG, used for randomly breaking ties\n",
    "        D (ndarray, ntaxa x ntaxa) : UPGMA d_ij or NJ D_ij\n",
    "\n",
    "    Returns:\n",
    "        (int) 0|1|2 : which pset tree is supported, based on min D_ij\n",
    "\n",
    "    D is either the unmodified d_ij distance matrix (UPGMA) or the\n",
    "    modified D_ij matrix for neighbor-joining.\n",
    "\n",
    "    Identify which pair i,j would be joined first, by finding\n",
    "    argmin_{ij} D_ij.\n",
    "\n",
    "    If more than one i,j pair has the same minimum cost, choose one \n",
    "    randomly.\n",
    "\n",
    "    From that first pair (i,j), identify which of the three pset tree\n",
    "    topologies would be the end result of the distance method. This\n",
    "    works because there are only three unrooted tree topologies for\n",
    "    4 taxa. \n",
    "\n",
    "    Rooting is ignored for this purpose. The three pset tree\n",
    "    topologies are treated as unrooted topologies. UPGMA finds rooted\n",
    "    topologies, and there are 15 possible rooted topologies for 4\n",
    "    taxa, but we only compare the UPGMA result against the pset trees\n",
    "    as unrooted topologies.\n",
    "    \"\"\"\n",
    "    (ntaxa,ntaxa) = D.shape\n",
    "    assert(ntaxa == 4)    # Not general; requires that we're using pset trees.\n",
    "\n",
    "    # Identify *a* argmin_ij D_ij.\n",
    "    # There may be more than one; if so, choose randomly.\n",
    "    #\n",
    "    # We can almost just do np.argmin(d) (and unravel that flattened\n",
    "    # idx) but we need to not count the diagonal where d_ii = 0, and\n",
    "    # we need to break any ties randomly.\n",
    "    #\n",
    "    # So: choose a random min i,j element from upper triangular d\n",
    "    # This is a carefully crafted NumPy incantation:\n",
    "    #    np.triu_indices gives us indices for upper triangular part of d_ij\n",
    "    #    np.min gives us the min value in it\n",
    "    #    np.flatnonzero gives array of one or more element indices in iu that have that min\n",
    "    #    rng.choice chooses one of them randomly, and sets that idx to k\n",
    "    #    best (i,j) is iu[0][k], iu[1][k]\n",
    "    # then we call best_tree to turn the first neighbor pair (i,j) into\n",
    "    # the index of the pset tree that that implies.\n",
    "    # \n",
    "    iu      = np.triu_indices(ntaxa,1)\n",
    "    k       = rng.choice(np.flatnonzero(D[iu] == np.min(D[iu])))\n",
    "    i,j     = iu[0][k],iu[1][k]\n",
    "\n",
    "    if   (i,j) == (0,1): best_tree = 1\n",
    "    elif (i,j) == (0,2): best_tree = 0\n",
    "    elif (i,j) == (0,3): best_tree = 2\n",
    "    elif (i,j) == (1,2): best_tree = 2\n",
    "    elif (i,j) == (1,3): best_tree = 0\n",
    "    elif (i,j) == (2,3): best_tree = 1\n",
    "    else:                exit(\"that can't happen\")\n",
    "    return best_tree"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "443d7391-bdbb-4e9a-a801-007717044de6",
   "metadata": {},
   "source": [
    "Now \"mini UPGMA\" simply means find the minimum distance $d_{ij}$, and then figure out which unrooted tree topology results from joining that (i,j) pair first."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 9,
   "id": "7b4bcdb1-07fd-4ba0-ae30-52d10aae9e6f",
   "metadata": {},
   "outputs": [],
   "source": [
    "def mini_upgma(rng, msa):\n",
    "    r\"\"\"Determine which pset tree would be chosen by UPGMA.\n",
    "\n",
    "    Args:\n",
    "        rng:                       numpy RNG. Used to break any ties.\n",
    "        msa (ndarray, ntaxa x L):  digital MSA for observed leaf seqs\n",
    "\n",
    "    Return:\n",
    "        (int) 0|1|2 : which pset tree is chosen by UPGMA\n",
    "\n",
    "    We only need to do the first iteration of UPGMA, choosing \n",
    "    min_{i,j} d_ij, to know which pset tree UPGMA finds.\n",
    "    \n",
    "    Typically called as `mini_upgma(rng, msa[:ntaxa])` since our\n",
    "    sampled MSAs include both leaf seqs and internal ancestors.\n",
    "\n",
    "    Assumes the 3 rooted pset trees are the only options for tree\n",
    "    topology, so ntaxa==4.\n",
    "    \"\"\"\n",
    "    ntaxa = len(msa)\n",
    "    d     = dist_matrix(msa)\n",
    "    return best_distance_tree(rng, d)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5ac66eb5-8c77-4e73-b98e-459d59c6726a",
   "metadata": {},
   "source": [
    "### 3. implement mini NJ"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9bf74d89-f9fc-4627-b6cc-8dda2aec6001",
   "metadata": {},
   "source": [
    "Now mini NJ is easy too. We just process the $d_{ij}$ matrix through NJ's adjusted $D_{ij} = d_{ij} - (r_i + r_j)$, with $r_i = \\frac{1}{n-2} \\sum_{k=0}^{n-1} d_{ik}$ for $n$ taxa."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 10,
   "id": "e641ba6a-b7b9-4205-8f94-47fed984b7a7",
   "metadata": {},
   "outputs": [],
   "source": [
    "def mini_nj(rng, msa):\n",
    "    r\"\"\"Determine which pset tree would be chosen by neighbor-joining.\n",
    "\n",
    "    Like mini_upgma() above, but with neighbor-joining.\n",
    "    \"\"\"\n",
    "    ntaxa = len(msa)\n",
    "\n",
    "    d  = dist_matrix(msa)\n",
    "    r  = np.sum(d, axis=1) / (ntaxa-2)   # rows or columns is same, d_ij is symmetric\n",
    "\n",
    "    # The neighbor-joining D matrix: D_ij = d_ij - (r_i + r_j)\n",
    "    # where r_i = \\sum_{k} d_{ik} for all taxa k\n",
    "    #\n",
    "    D  = np.zeros( d.shape )\n",
    "    for i in range(ntaxa):\n",
    "        for j in range(i+1,ntaxa):\n",
    "            D[i,j] = D[j,i] = d[i,j] - r[i] - r[j]\n",
    "\n",
    "    return best_distance_tree(rng, D)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "915f7e99-5892-4b3e-865c-814fa6b15a65",
   "metadata": {},
   "source": [
    "### 4. implement Fitch parsimony"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b10fc237-a7f4-457e-a2fc-6004580233c3",
   "metadata": {},
   "source": [
    "The trickiest bit of this (and ML) is just dealing with my tree data structure. For interior nodes $y = n..2n-2$, `T[y]` is a tuple of data for the children of node $y$; `T[y][0]` and `T[y][1]` are the two elements of that tuple, the data for left and right child (also a tuple); `T[y][w][0]` is the index of a child of $y$, where $w=0$ for left child and $w=1$ for right."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 11,
   "id": "ecc8b386-0c06-4ca1-9bf9-2c2cc54cab0b",
   "metadata": {},
   "outputs": [],
   "source": [
    "def parsimony(T, msa):\n",
    "    r\"\"\"Fitch parsimony: calculate cost of tree topology T for input MSA.\n",
    "\n",
    "    Args:\n",
    "        T (list of <nnodes> tree nodes): proposed tree topology\n",
    "        msa (ndarray, ntaxa x L):        ungapped digital MSA\n",
    "\n",
    "    Returns:\n",
    "        (int) cost of tree T.\n",
    "\n",
    "    The number of taxa and nodes in T are both determined from T.\n",
    "\n",
    "    msa can contain more aligned seqs (we generate MSA's that also\n",
    "    contain the ancestors, in rows ntaxa..nnodes-1, but they won't\n",
    "    be accessed. Only rows 0..ntaxa-1 of <msa> are accessed.\n",
    "    \"\"\"\n",
    "    nnodes = len(T)\n",
    "    ntaxa  = (nnodes + 1) // 2\n",
    "    ncol   = len(msa[0])\n",
    "\n",
    "    cost   = 0\n",
    "    S = [set() for _ in range(nnodes)]   # We'll construct a set of residues at each interior node...\n",
    "    for i in range(ncol):                # ... for each position i (aligned column) in seq alignment <msa>...\n",
    "        for y in range(ntaxa):\n",
    "            S[y] = { msa[y,i] }\n",
    "        for y in range(ntaxa, nnodes):\n",
    "            lc = T[y][0][0]\n",
    "            rc = T[y][1][0]\n",
    "            if len(S[lc] & S[rc]) > 0: S[y] = S[lc] & S[rc]\n",
    "            else:                      S[y] = S[lc] | S[rc]; cost += 1\n",
    "    return cost"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "95bd5388-fa9b-4f2a-8bb3-aba32636d991",
   "metadata": {},
   "source": [
    "### 5. implement Felsenstein maximum likelihood"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f54bb205-8840-4ef1-9cf0-5ab75b5ff0b3",
   "metadata": {},
   "source": [
    "The Felsenstein \"peeling\" algorithm is a dynamic programming algorithm that feels like a cross between the parsimony algorithm and the HMM forward algorithm that we've seen earlier in the course."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 12,
   "id": "faeace87-4250-44db-a749-6170e10f1809",
   "metadata": {},
   "outputs": [],
   "source": [
    "def maximum_likelihood(T, msa):\n",
    "    r\"\"\"Felstenstein ML algorithm: calculate log likelihood for tree T\n",
    "\n",
    "    Args:\n",
    "        T (list of <nnodes> tree nodes): proposed tree topology\n",
    "        msa (ndarray, ntaxa x L):        ungapped digital MSA\n",
    "\n",
    "    Returns:\n",
    "        (float) log likelihood of tree T. Larger is better.\n",
    "    \"\"\"\n",
    "    nnodes = len(T)\n",
    "    ntaxa  = (nnodes + 1) // 2\n",
    "    ncol   = len(msa[0])\n",
    "\n",
    "    # Precompute P(b|a,t) for each child branch\n",
    "    Prob = []\n",
    "    for y in range(ntaxa):         Prob.append(None)\n",
    "    for y in range(ntaxa, nnodes): Prob.append( (p_jukescantor(T[y][0][1]), p_jukescantor(T[y][1][1])) )  # left, right branch lens\n",
    "\n",
    "    logL = 0.0\n",
    "    # Sum -logL over each column in the alignment independently:\n",
    "    for i in range(ncol):\n",
    "        # Initialize with P(L_y | a) at leaves\n",
    "        L = np.zeros( (nnodes, 4) )        \n",
    "        for y in range(ntaxa):  L[y,msa[y,i]] = 1.0\n",
    "\n",
    "        for y in range(ntaxa, nnodes):\n",
    "            lc,rc = T[y][0][0], T[y][1][0]\n",
    "            lP,rP = Prob[y][0], Prob[y][1]\n",
    "            for a in range(4):\n",
    "                lprob,rprob = 0.0,0.0\n",
    "                for b in range(4):\n",
    "                    lprob += L[lc,b] * lP[a,b]\n",
    "                    rprob += L[rc,b] * rP[a,b]\n",
    "                L[y,a] = lprob * rprob\n",
    "\n",
    "        totprob = 0.\n",
    "        for a in range(4):\n",
    "            totprob += L[nnodes-1,a] * 0.25   # that's a \\pi_a = 0.25 uniform prior at the root\n",
    "        logL += np.log(totprob)               # summing log likelihood over positions/columns i\n",
    "    return logL\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "da851333-9915-450a-9833-236eea66b7c6",
   "metadata": {},
   "source": [
    "### 6. do Adler's suggested experiment"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "89a83dd3-b1c6-464e-9896-e97d0e40212e",
   "metadata": {},
   "source": [
    "Now we're ready to do the simulation experiment."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 13,
   "id": "48f0e41a-072b-4b86-83ce-1dccf9a27347",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "upgma:\n",
      "   len     T0     T1     T2\n",
      "    20      5     93      2\n",
      "   100      0    100      0\n",
      "   500      0    100      0\n",
      "  2000      0    100      0\n",
      "\n",
      "nj:\n",
      "   len     T0     T1     T2\n",
      "    20     50     33     17\n",
      "   100     77     17      6\n",
      "   500     97      3      0\n",
      "  2000    100      0      0\n",
      "\n",
      "parsimony:\n",
      "   len     T0     T1     T2\n",
      "    20     40     46     14\n",
      "   100     42     54      4\n",
      "   500     39     61      0\n",
      "  2000     26     74      0\n",
      "\n",
      "ml:\n",
      "   len     T0     T1     T2\n",
      "    20     58     23     19\n",
      "   100     83     11      6\n",
      "   500    100      0      0\n",
      "  2000    100      0      0\n",
      "\n"
     ]
    }
   ],
   "source": [
    "rng            = np.random.default_rng()\n",
    "nruns          = 100\n",
    "msalen_choices = [ 20, 100, 500, 2000]\n",
    "\n",
    "Trees  = specify_trees()\n",
    "ntrees = len(Trees)\n",
    "nnodes = len(Trees[0])\n",
    "ntaxa  = (nnodes + 1) // 2\n",
    "costs  = np.empty(ntrees)    # parsimony costs\n",
    "logL   = np.empty(ntrees)    # ML log likelihoods\n",
    "\n",
    "wins = {\n",
    "    'upgma'     : np.zeros( (len(msalen_choices), ntrees), dtype=int),\n",
    "    'nj'        : np.zeros( (len(msalen_choices), ntrees), dtype=int),\n",
    "    'parsimony' : np.zeros( (len(msalen_choices), ntrees), dtype=int),\n",
    "    'ml'        : np.zeros( (len(msalen_choices), ntrees), dtype=int),\n",
    "}\n",
    "\n",
    "for e, msalen in enumerate(msalen_choices):\n",
    "    for r in range(nruns):\n",
    "        msa  = sample_sequences(rng, Trees[0], msalen)\n",
    "\n",
    "        # UPGMA\n",
    "        best_tree = mini_upgma(rng, msa[:ntaxa])\n",
    "        wins['upgma'][e,best_tree] += 1\n",
    " \n",
    "        # NJ\n",
    "        best_tree = mini_nj(rng, msa[:ntaxa])\n",
    "        wins['nj'][e,best_tree] += 1\n",
    "            \n",
    "        # Parsimony\n",
    "        for t in range(ntrees):\n",
    "            costs[t]  = parsimony(Trees[t], msa[:ntaxa])\n",
    "        best_tree = rng.choice(np.flatnonzero(costs == np.min(costs)))\n",
    "        wins['parsimony'][e,best_tree] += 1\n",
    "\n",
    "        # ML\n",
    "        for t in range(ntrees):\n",
    "            logL[t]  = maximum_likelihood(Trees[t], msa[:ntaxa])\n",
    "        best_tree = rng.choice(np.flatnonzero(logL == np.max(logL)))\n",
    "        wins['ml'][e,best_tree] += 1\n",
    "\n",
    "for method in wins.keys():\n",
    "    print(f'{method}:')\n",
    "    print('{:>6s} {:>6s} {:>6s} {:>6s}'.format('len', 'T0', 'T1', 'T2'))\n",
    "    for e, msalen in enumerate(msalen_choices):\n",
    "        print('{:6d} {:6d} {:6d} {:6d}'.format(msalen, wins[method][e,0], wins[method][e,1], wins[method][e,2]))\n",
    "    print('')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cf3aad68-4716-40df-bfa2-ac625d0cc7d5",
   "metadata": {},
   "source": [
    "UPGMA systematically infers the wrong tree, T1. That's because the closest pair of sequences (by total summed branch length on the true tree T0) is the pair (0,1), which differ by 0.3 substitutions/site, but they aren't neighbors on the T0 tree topology. The true tree is additive but not ultrametric. It violates UPGMA's assumption that the tree is ultrametric (all taxa equidistant from the root). \n",
    "\n",
    "Parsimony, perhaps surprisingly, also systematically infers the wrong tree T1, but for a different reason. This artifact is called \"long branch attraction\". In the four-taxon MSA, the only \"phylogenetically informative\" sites (alignment columns) are those where two sequences have an identical residue, and the other two sequences have a different identical residue, in what I'll call a 2:2 pattern:\n",
    "\n",
    "* if there are no substitutions in any sequence (all four residues identical), all trees give a cost of 0\n",
    "* if there is 1 substitution in one sequence (3:1), all trees give a cost of 1\n",
    "* if the pattern is 2:2, then the tree that puts these pairs together has cost 1, and other trees have cost 2\n",
    "* if the pattern is 2:1:1, then all trees have cost 2 (not intuitive, but true)\n",
    "* if the pattern is 1:1:1:1, all trees have cost 3\n",
    "\n",
    "So if we happen to convergently generate the same substitution independently twice on the branches down to 2 and 3, without making any substitutions on the way to 0 and 1, that will look like support for tree T1 (which will explain those two substitutions as one event). The long branches to 2 and 3 increase the chances of this happening. \n",
    "\n",
    "This typically requires a fairly unusual tree, with unbalanced branch lengths that are long enough to have a high chance of convergent substitutions. But even a toy example like this (admittedly a carefully crafted one) can illustrate the artifact.\n",
    "\n",
    "ML is statistically consistent, meaning that if we generate data from the substitution probability model and use the same substitution model for inference, we will infer the correct tree, in the limit of infinite sites. We see it quickly converge to correctly inferring T0 here. Of course, we don't know the real substitution process for real biological data. Whether ML infers correct trees for real data is a different matter. We typically don't have the compute power to fully explore all trees, and our substitution models are oversimplified relative to real sequence evolution.\n",
    "\n",
    "Neighbor-joining works too! It sometimes gets a weird rap in the literature that I don't understand, perhaps because it's \"just\" a distance method and people lump it in with other distance methods like UPGMA. Neighbor-joining is statistically consistent and correctly infers the true tree, if the tree has additive distances. Additivity is not an unusual thing. When we use probabilistic models like Jukes-Cantor or Kimura or Hashegawa-Kishino-Yano to infer distances from observed alignments, we're inferring additive distances. "
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7fcf6961-e12b-4c19-be17-af70826e5647",
   "metadata": {},
   "source": [
    "### closing formalities"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "id": "93b61796-81a6-4a78-b50c-c7cb14823a55",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Python implementation: CPython\n",
      "Python version       : 3.13.1\n",
      "IPython version      : 8.30.0\n",
      "\n",
      "jupyter   : 1.1.1\n",
      "numpy     : 2.2.0\n",
      "matplotlib: 3.10.0\n",
      "\n",
      "Compiler    : Clang 15.0.0 (clang-1500.3.9.4)\n",
      "OS          : Darwin\n",
      "Release     : 25.5.0\n",
      "Machine     : arm64\n",
      "Processor   : arm\n",
      "CPU cores   : 14\n",
      "Architecture: 64bit\n",
      "\n"
     ]
    }
   ],
   "source": [
    "%load_ext watermark\n",
    "%watermark -v -m -p jupyter,numpy,matplotlib"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "30da910f-ab17-4f0d-87ad-706b5110f53c",
   "metadata": {},
   "source": [
    "[That's it!](https://en.wikipedia.org/wiki/So_Long,_and_Thanks_for_All_the_Fish)  Enjoy your summer. Thanks for all the work you put into MCB112."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3507999e-dd1d-4191-a8b6-1fb2b3f9dd6b",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.13.1"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
