{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c87d7a15-2d16-490a-8cb6-160a5ba16732",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import math\n",
    "import matplotlib.pyplot as plt\n",
    "import random\n",
    "import seaborn as sns\n",
    "from scipy.special import logsumexp\n",
    "\n",
    "base_to_idx = {'A':0, 'C':1, 'G':2, 'T':3}"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f9e2ca6e-d7c0-4003-a0e4-4060cdda1f49",
   "metadata": {},
   "source": [
    "# The case of the misappropriated reads"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "ca60fb0d-0421-4bc2-87f2-54688a5c1b41",
   "metadata": {},
   "source": [
    "Your lab has recently identified some new CRISPR systems and your PI has tasked divided them up between you, Alderman and Moriarty.  You have already confirmed they are active based on guides that target sequenced phages.  Now you want to do a PAM screen to determine the PAM motif.  This is a simple in vitro experiment involving a target plasmid library where your selected guide targeting region is flanked by random nucleotides PAM motif would be.  By expressing the CRISPR system only those target plasmids with compatible PAMs will be cleaved. By ligating adapters to the cleaved library, you can sequence just those targets that the CRISPR system cleaved with Illumina sequencing (https://www.youtube.com/watch?v=fCd6B5HRaZ8).  Illumina also allows you to pool many separate experiments together all at once by using different Barcodes to uniquely identify your samples.  By collecting many sequences you can create a PWM i.e. the consensus PAM motif.\n",
    "\n",
    "Alderman has already ran a trial experiment for his CRISPR system and gives you the reads see see what a PAM looks like (Alderman_reads).  This is super easy and something we have seen before and talked about in class, below are some functions to calculate a PWM from raw sequences and plotting functions.  Run the below cells to see the PAM!"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2184e127-3bdf-4614-a893-352f6178473c",
   "metadata": {
    "jupyter": {
     "source_hidden": true
    }
   },
   "outputs": [],
   "source": [
    "# Functions for plotting and such\n",
    "bases = ['A','C','G','T']\n",
    "colors = {'A':'green','C':'blue','G':'orange','T':'red'}\n",
    "\n",
    "def read_sequence(filename):\n",
    "    '''\n",
    "    Read a .txt file into a list of reads\n",
    "    '''\n",
    "    sequences = []\n",
    "    with open(filename) as file:\n",
    "        for line in file:\n",
    "            sequences.append(line[:-1]) # Prevent reading in the newline character\n",
    "    return np.array(sequences)\n",
    "\n",
    "def calculate_pwm(seq_list):\n",
    "    '''\n",
    "    Converts a list of DNA sequences into 4xn PAM pwm \n",
    "    '''\n",
    "    n_positions = len(seq_list[0])\n",
    "    pam_pwm = np.zeros((4, n_positions), dtype=float)\n",
    "\n",
    "    for seq in seq_list:\n",
    "        for position, base in enumerate(seq):\n",
    "            pam_pwm[base_to_idx[base], position] += 1 # rows are the frequencies A|C|G|T and columns are positions\n",
    "\n",
    "    return pam_pwm/len(seq_list)\n",
    "\n",
    "def pwm_to_information(pwm):\n",
    "    '''\n",
    "    Converts a PWM into 2-Bit information for plotting\n",
    "    '''\n",
    "    info_matrix = np.zeros_like(pwm)\n",
    "    n_positions = pwm.shape[1]\n",
    "\n",
    "    for i in range(n_positions):\n",
    "        col = pwm[:, i]\n",
    "        entropy = -np.sum([p*np.log2(p) if p > 0 else 0 for p in col]) # column entropy (skip zeros)\n",
    "        R = 2 - entropy # Calculate information, max for DNA = 2 bits\n",
    "        info_matrix[:, i] = col * R # letter heights\n",
    "\n",
    "    return info_matrix\n",
    "\n",
    "def plot_information_logo(info_matrix, title):\n",
    "    '''\n",
    "    Plots the PAM information matrix as a WEBlogo like plot\n",
    "    '''\n",
    "    n_positions = info_matrix.shape[1]\n",
    "    fig, ax = plt.subplots(figsize=(6,3))\n",
    "\n",
    "    for i in range(n_positions):\n",
    "        bottom = 0\n",
    "        col_heights = info_matrix[:, i]\n",
    "        sorted_idx = np.argsort(col_heights) # sort letters by height ascending (smallest at bottom)\n",
    "        for idx in sorted_idx:\n",
    "            height = col_heights[idx]\n",
    "            base = bases[idx]\n",
    "            if height > 0:\n",
    "                ax.bar(i, height, bottom=bottom, color=colors[base], width=0.8) # Draw the colored rectangle\n",
    "                if height > 0.2: # Draw the letter inside\n",
    "                    ax.text(i, bottom + height/2, base, ha='center', va='center',\n",
    "                            fontsize=14, fontweight='bold', color='white')\n",
    "                bottom += height  # stack next letter on top\n",
    "    ax.set_xlim(-0.5, n_positions - 0.5)\n",
    "    ax.set_ylim(0, 2)\n",
    "    ax.set_xticks(range(n_positions))\n",
    "    ax.set_xlabel(\"Position\")\n",
    "    ax.set_ylabel(\"Information (bits)\")\n",
    "    ax.set_title(title)\n",
    "    sns.despine()\n",
    "    plt.tight_layout()\n",
    "    plt.show()\n",
    "\n",
    "def plot_frequencies(pwm):\n",
    "    '''\n",
    "    Plots the raw PWM frequencies - easier to see the toggle bases for PAMs\n",
    "    '''\n",
    "    n_positions = pwm.shape[1]\n",
    "    fig, ax = plt.subplots(figsize=(6,3))\n",
    "\n",
    "    for i in range(n_positions):\n",
    "        bottom = 0\n",
    "        col_heights = pwm[:, i]\n",
    "        sorted_idx = np.argsort(col_heights) # sort letters by height ascending (smallest at bottom)\n",
    "        for idx in sorted_idx:\n",
    "            height = col_heights[idx]\n",
    "            base = bases[idx]\n",
    "            if height > 0:\n",
    "                ax.bar(i, height, bottom=bottom, color=colors[base], width=0.8) # Draw the colored rectangle\n",
    "                if height > 0.2: # Draw the letter inside\n",
    "                    ax.text(i, bottom + height/2, base, ha='center', va='center',\n",
    "                            fontsize=14, fontweight='bold', color='white')\n",
    "                bottom += height  # stack next letter on top\n",
    "\n",
    "    ax.set_xlim(-0.5, n_positions - 0.5)\n",
    "    ax.set_ylim(0, 1)\n",
    "    ax.set_xticks(range(n_positions))\n",
    "    ax.set_xlabel(\"Position\")\n",
    "    ax.set_ylabel(\"Base Frequency\")\n",
    "    sns.despine()\n",
    "    plt.tight_layout()\n",
    "    plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "73f04700-dfaa-481f-b0d3-141e682ff94a",
   "metadata": {},
   "source": [
    "# Example PAM Screen"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "34bb2de1-2a50-4f99-a407-735d4c9643c3",
   "metadata": {},
   "outputs": [],
   "source": [
    "example_reads = read_sequence('Alderman_reads.txt')\n",
    "example_PAM = calculate_pwm(example_reads)\n",
    "plot_information_logo(pwm_to_information(example_PAM), 'Alderman PAM Screen')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b81a31b4-29d7-424a-ac17-dcf134d410fc",
   "metadata": {},
   "source": [
    "Nice! You have succesfully derived the PAM T T/A T A C.  You are now prepared to run the experiment yourself.  You perform the PAM screen experiment and prepare your library for sequencing. Since you are using Illumina Sequencing Moriarty and Alderman ask if they can pool their samples with you (Alderman wants to get some more reads to be super sure).  Afterwards you demultiplex the sequencing to get your sample sequences and you run the code and get the following PAM:"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cf6e8cbd-d7d7-4a68-b7ac-179f535585c1",
   "metadata": {},
   "source": [
    "# Your Data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6fbbb6dd-9aea-4a82-938c-4250c8a2d3f0",
   "metadata": {},
   "outputs": [],
   "source": [
    "your_reads = read_sequence('Your_reads.txt')\n",
    "your_PAM = calculate_pwm(your_reads)\n",
    "\n",
    "plot_information_logo(pwm_to_information(your_PAM), 'Your PAM Screen')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ad2ae3b0-fe48-498f-9b03-db491414c204",
   "metadata": {},
   "outputs": [],
   "source": [
    "moriarty_reads = read_sequence('Moriarty_reads.txt')\n",
    "moriarty_PAM = calculate_pwm(moriarty_reads)\n",
    "\n",
    "plot_information_logo(pwm_to_information(moriarty_PAM), 'Moriarty PAM Screen')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "af129ea7-5fc6-4d18-978f-2ed9a39288e3",
   "metadata": {},
   "source": [
    "This looks terrible! Clearly something went wrong! Later Moriarty approaches you and asks if your sequencing worked because he also got garbage for his PAM (check his as well), but it does seem like there is the hint of a motif in his.  Alderman blames you thinking your library preparation or loading of the sequencer must have been wrong and storms off to tell your PI. You notice that the actual reads look fine so you ask Moriarty what Illumina primers he used and and he says he got them from Alderman (who's handwriting is notoriously bad).  When he shows you the plate you notice that the 1s and 7s look awfully similar. You think that maybe the same sets of barcodes got used for all the experiments, dumping all the sequencing into the same sample files!  You are about to despair but remember you just learned about Expectation Maximization in MCB112 (sometimes known as a mixture models) and this seems awfully similar to the scenario described in class for finding the RBS motif, maybe even simpler.  You decide to try and deconvolute the reads before Alderman can find your PI!\n"
   ]
  },
  {
   "attachments": {},
   "cell_type": "markdown",
   "id": "0e7ba621-962b-4a6f-86f0-c78a7712f059",
   "metadata": {},
   "source": [
    "In the scenario discussed in class our hidden variable $\\lambda$ is the start positions of the RBS motif in each sequence.  For the PAM screen experiment we know exactly where the motif should be in the sequence, but because of the barcode mix up we don't know to which sample (and motif) the sequence belongs to. So now our hidden variable $\\lambda$ is the \"responsibility\" of each sequence i.e. the probability the sequence came from a particular PAM.  Unlike the homework we don't have to worry about start position and a background model, but we do have to keep track of multiple PWMs now.  Similarly I will summarize the model which is a the collection of all the possible PAMs as $\\theta$, and a superscript $m$ specifies the specific PAM:\n",
    "\n",
    "$$\n",
    "  p_k^m(a) = \\frac{c_k^m(a)} {\\sum_b c_k^m(b)}\n",
    "$$\n",
    "\n",
    "\n",
    "$$\n",
    "   P(\\lambda^m \\mid \\theta, X^s) = \n",
    "     \\frac{P(X^s \\mid \\theta, \\lambda^m) P(\\lambda^m)}\n",
    "          {\\sum_{\\lambda'} P(X^s \\mid \\theta, \\lambda') P(\\lambda')}\n",
    "$$\n",
    "\n",
    "$$\n",
    "   P(X^s \\mid \\theta, \\lambda^m) =\n",
    "     \\prod_{k=1}^{W} p_k^m(X^s_{k}) \n",
    "$$\n",
    "\n",
    "$$\n",
    "  c_k^m(a) = \\sum_{s=1}^N P(\\lambda^m \\mid \\theta, X^s)\\; \\delta(X^s_{k} = a)\n",
    "$$\n",
    "\n",
    "$$\n",
    "\\hat{p}_i^{m} = \\frac{c_i^{m} + \\alpha_i}{\\sum_j \\left(c_j^{m} + \\alpha_j\\right)}\n",
    "$$\n",
    "\n",
    "This at least looks a lot more simple than what the math in class was! A difference not seen in the math here is that we have to keep track of each pam $m$ somewhere but that is easily done in a list.\n",
    "\n",
    "Now $P(\\lambda^m \\mid \\theta)$ is the prior which in this case is the underlying frequency of each PAM i.e. how much of each sequence comes from what underlying PAM.  In class this term was the prior of the starting position for the motif which we assumed to be uniform, thus being a constant (meaning we can omit it in the actual EM since it cancels out). We could in principal figure this out as well during our EM because it is not necessarily uniform (different amounts of reads from each experiments).  However we can also assume that we loaded the library perfectally and an even amout of reads were attributed to each sample and therefore the prior is evenly distributed i.e. $1/M$ where M is the number of PAMs, which is more similar to the pset so we will do that.  In which case we can ignore it in the actual EM because it is similarly a constant.\n",
    "\n",
    "To implements, we would repeat all of these steps for each PAM $m$ until reaching some convergence.  Hopefully I wrote this all down correctly, if you find an error let me know and you will get a prize!\n",
    "\n",
    "What is core to this if you can't already tell is calculating $P(X^s \\mid \\theta, \\lambda^m)$ a lot. So we will probably want to write a function that can do this just on its own.  Also as a reminder base_to_idx = {'A':0, 'C':1, 'G':2, 'T':3} will be used a bunch to to make sure everything stays consistent."
   ]
  },
  {
   "cell_type": "raw",
   "id": "589fa473-1e57-4b5d-9f39-0e33fe70be8d",
   "metadata": {},
   "source": [
    "Spot to write some pseudocode notes if you care"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5c03bbb9-8d9a-47b6-ad33-24e715222b59",
   "metadata": {},
   "source": [
    "# Calculate Likelihood Bank"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ed97042c-c168-4b39-8c78-537ba2918e4a",
   "metadata": {},
   "outputs": [],
   "source": [
    "def calc_likelihood_1(seq, pwm):\n",
    "    log_prob = 0.0\n",
    "    \n",
    "    for pos, base in enumerate(seq):\n",
    "        base_index = base_to_idx[base]\n",
    "        prob = pwm[pos, base_index]\n",
    "        log_prob += np.log2(prob)\n",
    "    \n",
    "    return log_prob\n",
    "\n",
    "\n",
    "def calc_likelihood_2(seq, pwm, pseudocount=1e-6):\n",
    "    log_prob = 0.0\n",
    "    \n",
    "    for pos, base in enumerate(seq):\n",
    "        base_index = base_to_idx[base]\n",
    "        prob = pwm[pos, base_index]\n",
    "        log_prob += np.log2(prob + pseudocount)\n",
    "    \n",
    "    return log_prob\n",
    "\n",
    "\n",
    "def calc_likelihood_3(seq, pwm, pseudocount=1e-6):\n",
    "    prob = 1.0\n",
    "    \n",
    "    for pos, base in enumerate(seq):\n",
    "        base_index = base_to_idx[base]\n",
    "        prob *= (pwm[pos, base_index] + pseudocount)\n",
    "    \n",
    "    return np.log2(prob)\n",
    "\n",
    "\n",
    "def calc_likelihood_4(seq, pwm, pseudocount=1e-6):\n",
    "    log_prob = 0.0\n",
    "    for pos, base in enumerate(seq):\n",
    "        base_index = base_to_idx[base]\n",
    "        log_prob += np.log2(pwm[pos, base_index]) + pseudocount \n",
    "        \n",
    "    return log_prob\n",
    "\n",
    "\n",
    "    "
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7a235fe5-a89c-4865-8bd2-b3c7bd327802",
   "metadata": {},
   "source": [
    "# Initialization Bank"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e656db18-5ad0-464a-aa1c-e0d8c31f2eb2",
   "metadata": {},
   "outputs": [],
   "source": [
    "def initialize_pwms_1(num_pams, sequences):\n",
    "    W = len(sequences[0]) # Find the width of the PAM\n",
    "    \n",
    "    pwm = np.full((W, 4), 0.25) # Create PAM of equal probs \n",
    "    pwms = np.tile(pwm[None, :, :], (num_pams, 1, 1)) # Duplicate \n",
    "    \n",
    "    return pwms\n",
    "\n",
    "\n",
    "def initialize_pwms_2(num_pams, sequences):\n",
    "    W = len(sequences[0]) # Find the width of the PAM\n",
    "    \n",
    "    pwm = np.random.rand(W, 4) # Create one i.i.d. PAM             \n",
    "    pwm /= pwm.sum(axis=1, keepdims=True) # Normalize\n",
    "    pwms = np.repeat(pwm[np.newaxis, :, :], num_pams, axis=0) # Duplicate \n",
    "    \n",
    "    return pwms\n",
    "\n",
    "\n",
    "def initialize_pwms_3(num_pams, sequences):\n",
    "    W = len(sequences[0]) # Find the width of the PAM\n",
    "    \n",
    "    pwms = np.random.rand(num_pams, W, 4) # Create several i.i.d. PAMs             \n",
    "    pwms /= pwms.sum(axis=2, keepdims=True) # Normalize      \n",
    "    \n",
    "    return pwms\n",
    "\n",
    "\n",
    "def initialize_pwms_4(num_pams, sequences, pseudocount=1e-6):\n",
    "    W = len(sequences[0]) # Find the width of the PAM\n",
    "    \n",
    "    counts = np.full((num_pams, W, 4), fill_value=pseudocount)\n",
    "    for seq in sequences: # Assign each sequence ranomly to PAM\n",
    "        k = np.random.randint(0, num_pams)\n",
    "        for i, base in enumerate(seq):\n",
    "            counts[k, i, base_to_idx[base]] += 1 # Count base\n",
    "    # Determine PAM from raw counts\n",
    "    pwms = counts / counts.sum(axis=2, keepdims=True) \n",
    "    \n",
    "    return pwms\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "4166f917-24d3-4840-ab52-741ccb17edeb",
   "metadata": {},
   "source": [
    "# E-step Bank"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "12039607-51ed-4f74-8c54-5c08c14f843d",
   "metadata": {},
   "outputs": [],
   "source": [
    "# I know it is clunky to have to pass calc_likelihood into another function \n",
    "# but we have to for this whole option game to work so sorry\n",
    "# Also to make the functions easier to read I have not included the expected counts in the E-step\n",
    "# I find that conceptually easier to think of in the traditional M-Step\n",
    "\n",
    "def E_step_1(calc_likelihood, pam_pwms, sequences):\n",
    "    num_pams = len(pam_pwms)\n",
    "    post_matrix = np.zeros((len(sequences), num_pams))\n",
    "    \n",
    "    for i, seq in enumerate(sequences): # Loop through each sequence\n",
    "        log_likelihoods = np.zeros(num_pams)\n",
    "        for s in range(num_pams): # Calculate LL for each PAM\n",
    "            log_likelihoods[s] = calc_likelihood(seq, pam_pwms[s])\n",
    "        post_matrix[i, :] = np.exp2(log_likelihoods) # Take out of Log Space\n",
    "\n",
    "    return post_matrix\n",
    "\n",
    "\n",
    "def E_step_2(calc_likelihood, pam_pwms, sequences):\n",
    "    num_pams = len(pam_pwms)\n",
    "    post_matrix = np.zeros((len(sequences), num_pams))\n",
    "    \n",
    "    for i, seq in enumerate(sequences): # Loop through each sequence\n",
    "        log_likelihoods = np.array([calc_likelihood(seq, pwm) \n",
    "                            for pwm in pam_pwms]) # Calculate LL for each PAM\n",
    "        max_idx = np.argmax(log_likelihoods) # Find most likely PAM\n",
    "        post_matrix[i, max_idx] = 1.0  # Assign read to that PAM\n",
    "\n",
    "    return post_matrix\n",
    "\n",
    "\n",
    "def E_step_3(calc_likelihood, pam_pwms, sequences):\n",
    "    num_pams = len(pam_pwms)\n",
    "    post_matrix = np.zeros((len(sequences), num_pams))\n",
    "\n",
    "    for i, seq in enumerate(sequences): # Loop through each sequence\n",
    "        log_likelihoods = np.zeros(num_pams)\n",
    "        for s in range(num_pams): # Calculate LL for each PAM \n",
    "            log_likelihoods[s] = calc_likelihood(seq, pam_pwms[s])\n",
    "        # Marginalize\n",
    "        post_matrix[i, :] = np.exp2(log_likelihoods - logsumexp(log_likelihoods)) \n",
    "\n",
    "    return post_matrix\n",
    "\n",
    "\n",
    "def E_step_4(calc_likelihood, pam_pwms, sequences):\n",
    "    num_pams = len(pam_pwms)\n",
    "    post_matrix = np.zeros((len(sequences), num_pams))\n",
    "\n",
    "    for i, seq in enumerate(sequences):  # Loop through each sequence\n",
    "        log_likelihoods = np.zeros(num_pams)\n",
    "        for s in range(num_pams):# Calculate LL for each PAM\n",
    "            log_likelihoods[s] = calc_likelihood(seq, pam_pwms[s])\n",
    "        probs = np.exp2(log_likelihoods) # Take out of log space\n",
    "        # Marginzalize\n",
    "        post_matrix[i, :] = probs / np.sum(probs)\n",
    "\n",
    "    return post_matrix\n",
    "\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "10a1d766-df9c-4b75-a8ad-986af864edb5",
   "metadata": {},
   "source": [
    "# M-step Bank"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "abebc7d8-fd9a-4e60-8af5-503cd21206c6",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Note I start the M-steps with calculating the expected counts, which we traditionally define as part of the E-step\n",
    "# However to make this easier to read I find it more intuitive to group the expected counts with updating the PWM\n",
    "# Otherwise the M-step is super trivial.  This isn't uncommon in my experience to group the counts in the M-step\n",
    "# But I do apologize if this makes it confusing!\n",
    "\n",
    "def M_step_1(post_matrix, sequences, pam_pwms, W=8):\n",
    "    num_pams = len(pam_pwms)\n",
    "    counts = np.zeros_like(pam_pwms)\n",
    "    new_pwms = np.zeros_like(pam_pwms)\n",
    "\n",
    "    for n in range(num_pams): # foreach PAM\n",
    "        for i, seq in enumerate(sequences): # Loop through each sequence\n",
    "            for pos in range(W): # Loop through PAM position\n",
    "                base_index = base_to_idx[seq[pos]] # convert base to index\n",
    "                counts[n][pos, base_index] += post_matrix[i, n] # count base\n",
    "        new_pwms[n] = (counts[n]) / np.sum(counts[n], axis=1, keepdims=True) # Normalize\n",
    "\n",
    "    return new_pwms\n",
    "\n",
    "    \n",
    "def M_step_2(post_matrix, sequences, pam_pwms, W=8, pseudocount=1):\n",
    "    num_pams = len(pam_pwms)\n",
    "    counts = np.zeros_like(pam_pwms)\n",
    "    new_pwms = np.zeros_like(pam_pwms)\n",
    "\n",
    "    for i, seq in enumerate(sequences): # Loop through sequences\n",
    "        k = np.argmax(post_matrix[i]) # Identify most likely PAM\n",
    "        for pos in range(W): # Loop through position\n",
    "            base_index = base_to_idx[seq[pos]] # convert base to index\n",
    "            counts[k][pos, base_index] += 1 # count base\n",
    "    for n in range(num_pams):\n",
    "        new_pwms[n] = (counts[n] + pseudocount) / np.sum(\n",
    "            counts[n] + pseudocount, axis=1, keepdims=True) # Normalize \n",
    "\n",
    "    return new_pwms\n",
    "\n",
    "\n",
    "def M_step_3(post_matrix, sequences, pam_pwms, W=8, pseudocount=1):\n",
    "    num_pams = len(pam_pwms)\n",
    "    counts = np.zeros_like(pam_pwms)\n",
    "    new_pwms = np.zeros_like(pam_pwms)\n",
    "\n",
    "    for n in range(num_pams): # loop through PAMs\n",
    "        for i, seq in enumerate(sequences): # Loop through sequence\n",
    "            for pos in range(W): # Loop through position\n",
    "                base_index = base_to_idx[seq[pos]] # convert base to index\n",
    "                counts[n][pos, base_index] += post_matrix[i, n] # count base\n",
    "        new_pwms[n] = (counts[n] + pseudocount) / np.sum(\n",
    "            counts[n, :, :] + pseudocount) # Normalize over total counts\n",
    "\n",
    "    return new_pwms\n",
    "\n",
    "\n",
    "def M_step_4(post_matrix, sequences, pam_pwms, W=8, pseudocount=1):\n",
    "    num_pams = len(pam_pwms)\n",
    "    counts = np.zeros_like(pam_pwms)\n",
    "    new_pwms = np.zeros_like(pam_pwms)\n",
    "\n",
    "    for n in range(num_pams): # For each PAM\n",
    "        for i, seq in enumerate(sequences): # Loop through each sequence\n",
    "            for pos in range(W): # Loop through each position\n",
    "                base_index = base_to_idx[seq[pos]]\n",
    "                counts[n][pos, base_index] += post_matrix[i, n] # Count bases\n",
    "        new_pwms[n] = (counts[n] + pseudocount) / np.sum(\n",
    "            counts[n] + pseudocount, axis=1, keepdims=True) # Normalize \n",
    "\n",
    "    return new_pwms\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "942c6712-9550-4822-8420-7c422766d4a5",
   "metadata": {},
   "source": [
    "# Implementation Bank"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "16d12b72-b125-4a77-9b5c-096578b4bd11",
   "metadata": {},
   "outputs": [],
   "source": [
    "def EM_1(calc_likelihood, initialize_pwms, E_step, M_step, X, num_PAMs, max_iters=25):\n",
    "    \n",
    "    PAMs = initialize_pwms(num_PAMs, X) # Initialize\n",
    "    for iteration in range(max_iters): # Iterate several times\n",
    "        posterior_matrix = E_step(calc_likelihood, PAMs, X) # E-step\n",
    "        PAMs = M_step(posterior_matrix, X, PAMs) # M-step\n",
    "        \n",
    "    return PAMs\n",
    "\n",
    "    \n",
    "def EM_2(calc_likelihood, initialize_pwms, E_step, M_step, X, num_PAMs, max_iters=50, epsilon=1e-4):\n",
    "    \n",
    "    PAMs = initialize_pwms(num_PAMs, X) # Initialize\n",
    "    for iteration in range(max_iters):\n",
    "        posterior_matrix = E_step(calc_likelihood, PAMs, X) # E-step\n",
    "        new_PAMs = M_step(posterior_matrix, X, PAMs) # M-step\n",
    "        \n",
    "        LL = 0 # Compute log-likelihood of each sequence\n",
    "        for seq in X:\n",
    "            log_probs = np.array([calc_likelihood(seq, pwm) for pwm in new_PAMs])\n",
    "            LL += logsumexp(log_probs)\n",
    "\n",
    "        if iteration > 1: # Convergence check \n",
    "            if abs( (LL_old - LL) / LL_old) < epsilon: # If the fold change is less than epsilon\n",
    "                break # End EM\n",
    "        if iteration > max_iters: # End if we time out\n",
    "            break\n",
    "        LL_old = LL # Update for next iteration\n",
    "        PAMs = new_PAMs\n",
    "        \n",
    "    return PAMs\n",
    "\n",
    "\n",
    "def EM_3(calc_likelihood, initialize_pwms, E_step, M_step, X, num_PAMs, epsilon=1e-4):\n",
    "    \n",
    "    PAMs = initialize_pwms(num_PAMs, X) # Initialize\n",
    "    LL_old = -np.inf # Starting LL      \n",
    "    \n",
    "    while True: # Keep on going\n",
    "        posterior_matrix = E_step(calc_likelihood, PAMs, X) # E-step\n",
    "        new_PAMs = M_step(posterior_matrix, X, PAMs) # M-step\n",
    "        \n",
    "        LL = 0 # Compute log-likelihood of each sequence\n",
    "        for seq in X:\n",
    "            log_probs = np.array([calc_likelihood(seq, pwm) for pwm in new_PAMs])\n",
    "            LL += logsumexp(log_probs)\n",
    "\n",
    "        if LL_old != -np.inf: # Convergence check\n",
    "            if abs((LL - LL_old) / abs(LL_old)) < epsilon: # If fold change is small\n",
    "                break # End EM\n",
    "        LL_old = LL # Update for next iteration\n",
    "        PAMs = new_PAMs # PAMs will be your final \n",
    "        \n",
    "    return PAMs\n",
    "\n",
    "\n",
    "def EM_4(calc_likelihood, initialize_pwms, E_step, M_step, X, num_PAMs, max_iters=25, trials=5):\n",
    "    \n",
    "    best_LL = -np.inf\n",
    "    \n",
    "    for trial in range(trials): # Attempt several tries\n",
    "        PAMs = initialize_pwms(num_PAMs, X) # initialize per trial\n",
    "        for iteration in range(max_iters): # Iterate several times\n",
    "            posterior_matrix = E_step(calc_likelihood, PAMs, X) # E-step\n",
    "            PAMs = M_step(posterior_matrix, X, PAMs) # M-step\n",
    "            \n",
    "        LL = 0 # Compute log-likelihood for trial\n",
    "        for seq in X:\n",
    "            log_probs = np.array([calc_likelihood(seq, pwm) for pwm in PAMs])\n",
    "            LL += logsumexp(log_probs)\n",
    "        # Track best trial\n",
    "        if LL > best_LL: # If new LL is better choose those PAMs\n",
    "            best_PAMs = PAMs.copy() # Update your best PAMs\n",
    "            best_LL = LL\n",
    "\n",
    "    return best_PAMs\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f1c52827-49a8-423c-aae5-e597fab65072",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Input your reads\n",
    "reads = read_sequence('')\n",
    "OG_PAM = calculate_pwm(reads)\n",
    "\n",
    "# Plot original PAM\n",
    "plot_information_logo(pwm_to_information(OG_PAM), 'Original PAM')\n",
    "\n",
    "# Perform EM!\n",
    "# edit below to choose your functions\n",
    "PAMs = EM_***(calc_likelihood_***, initialize_pwms_***, E_step_***, M_step_***, reads, ***)\n",
    "for i, pam in enumerate(PAMs, start=1): # Plot the output PAMs from the EM\n",
    "    plot_information_logo(pwm_to_information(pam.T), f\"PAM {i}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "486825ca-45c0-4700-9370-eb12f490f219",
   "metadata": {},
   "source": [
    "Does it looks like you expect it too?  What is your conclusion? What happens when we change the number of input PAMs? Make sure to try running the algorithm a few times (re local optimization and such) and try some different options for the functions!\n",
    "\n",
    "Also you have low-key generated a de-noising algorithm, try running this on the Noisy_reads.txt!"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "seq_analysis",
   "language": "python",
   "name": "seq_analysis"
  },
  "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.12.13"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
