{ "cells": [ { "cell_type": "markdown", "metadata": { "id": "view-in-github", "colab_type": "text" }, "source": [ "\"Open" ] }, { "cell_type": "markdown", "metadata": { "id": "OA2k3sAYuiXe" }, "source": [ "#af_relax_design (WIP)\n", "\n", "\n", "**Efficient and scalable de novo protein design using a relaxed sequence space**\n", "\n", "Christopher Josef Frank, Ali Khoshouei, Yosta de Stigter, Dominik Schiewitz, Shihao Feng, Sergey Ovchinnikov, Hendrik Dietz\n", "\n", "doi: https://doi.org/10.1101/2023.02.24.529906\n", "\n", "**WARNING** This notebook is in development, we are still working on adding all the options from the manuscript above." ] }, { "cell_type": "code", "execution_count": null, "metadata": { "cellView": "form", "id": "-AXy0s_4cKaK" }, "outputs": [], "source": [ "#@title setup\n", "%%time\n", "import os\n", "if not os.path.isdir(\"params\"):\n", " # get code\n", " os.system(\"pip -q install pyppeteer nest_asyncio\")\n", " os.system(\"pip -q install git+https://github.com/sokrypton/ColabDesign.git\")\n", " # for debugging\n", " os.system(\"ln -s /usr/local/lib/python3.*/dist-packages/colabdesign colabdesign\")\n", " # download params\n", " os.system(\"mkdir params\")\n", " os.system(\"apt-get install aria2 -qq\")\n", " os.system(\"aria2c -q -x 16 https://storage.googleapis.com/alphafold/alphafold_params_2022-12-06.tar\")\n", " os.system(\"tar -xf alphafold_params_2022-12-06.tar -C params\")\n", "\n", "import warnings\n", "warnings.simplefilter(action='ignore', category=FutureWarning)\n", "\n", "import os\n", "from colabdesign import mk_afdesign_model, clear_mem\n", "from colabdesign.mpnn import mk_mpnn_model\n", "\n", "from IPython.display import HTML\n", "from google.colab import files\n", "import numpy as np\n", "\n", "import requests, time\n", "if not os.path.isfile(\"TMscore\"):\n", " os.system(\"wget -qnc https://zhanggroup.org/TM-score/TMscore.cpp\")\n", " os.system(\"g++ -static -O3 -ffast-math -lm -o TMscore TMscore.cpp\")\n", "def tmscore(x,y):\n", " # pass to TMscore\n", " output = os.popen(f'./TMscore {x} {y}')\n", " # parse outputs\n", " parse_float = lambda x: float(x.split(\"=\")[1].split()[0])\n", " o = {}\n", " for line in output:\n", " line = line.rstrip()\n", " if line.startswith(\"RMSD\"): o[\"rms\"] = parse_float(line)\n", " if line.startswith(\"TM-score\"): o[\"tms\"] = parse_float(line)\n", " if line.startswith(\"GDT-TS-score\"): o[\"gdt\"] = parse_float(line)\n", " return o\n", "\n", "import asyncio\n", "import nest_asyncio\n", "from pyppeteer import launch\n", "import base64\n", "\n", "# Apply nest_asyncio to enable nested event loops\n", "nest_asyncio.apply()\n", "\n", "async def fetch_blob_content(page, blob_url):\n", " blob_to_base64 = \"\"\"\n", " async (blobUrl) => {\n", " const blob = await fetch(blobUrl).then(r => r.blob());\n", " return new Promise((resolve) => {\n", " const reader = new FileReader();\n", " reader.onloadend = () => resolve(reader.result);\n", " reader.readAsDataURL(blob);\n", " });\n", " }\n", " \"\"\"\n", " base64_data = await page.evaluate(blob_to_base64, blob_url)\n", " _, encoded = base64_data.split(',', 1)\n", " return base64.b64decode(encoded)\n", "\n", "async def extract_pdb_file_download_link_and_content(url):\n", " browser = await launch(headless=True, args=['--no-sandbox', '--disable-setuid-sandbox'])\n", " page = await browser.newPage()\n", " await page.goto(url, {'waitUntil': 'networkidle0'})\n", " elements = await page.querySelectorAll('a.btn.bg-purple')\n", " for element in elements:\n", " href = await page.evaluate('(element) => element.getAttribute(\"href\")', element)\n", " if 'blob:https://esmatlas.com/' in href:\n", " content = await fetch_blob_content(page, href)\n", " await browser.close()\n", " return href, content\n", " await browser.close()\n", " return \"No PDB file link found.\", None\n", "\n", "def esmfold_api(sequence):\n", " url = f'https://esmatlas.com/resources/fold/result?fasta_header=%3Eunnamed&sequence={sequence}'\n", " result = asyncio.get_event_loop().run_until_complete(extract_pdb_file_download_link_and_content(url))\n", " if result[1]:\n", " pdb_str = result[1].decode('utf-8')\n", " return pdb_str\n", " else:\n", " return \"Failed to retrieve PDB content.\"\n", "\n", "import jax\n", "import jax.numpy as jnp\n", "from colabdesign.af.alphafold.common import residue_constants" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "sZnYfCbfEvol", "cellView": "form" }, "outputs": [], "source": [ "#@title # hallucination\n", "#@markdown For a given length, generate/hallucinate a protein sequence that AlphaFold thinks folds into a well structured protein (high plddt, low pae, many contacts).\n", "LENGTH = 100 #@param {type:\"integer\"}\n", "COPIES = 1 #@param [\"1\", \"2\", \"3\", \"4\", \"5\", \"6\", \"7\", \"8\"] {type:\"raw\"}\n", "MODE = \"manuscript\" #@param [\"original\", \"manuscript\"]\n", "use_rg_loss = True #@param {type:\"boolean\"}\n", "\n", "#@markdown ProteinMPNN Settings\n", "use_mpnn_loss = False #@param {type:\"boolean\"}\n", "use_solubleMPNN = False #@param {type:\"boolean\"}\n", "#@markdown\n", "\n", "def add_rg_loss(self, weight=0.1):\n", " '''add radius of gyration loss'''\n", " def loss_fn(inputs, outputs):\n", " xyz = outputs[\"structure_module\"]\n", " ca = xyz[\"final_atom_positions\"][:,residue_constants.atom_order[\"CA\"]]\n", " if self.protocol == \"binder\":\n", " ca = ca[-self._binder_len:]\n", " if MODE == \"manuscript\":\n", " ca = ca[::5]\n", " rg = jnp.sqrt(jnp.square(ca - ca.mean(0)).sum(-1).mean() + 1e-8)\n", " if MODE == \"original\":\n", " rg_th = 2.38 * ca.shape[0] ** 0.365\n", " rg = jax.nn.elu(rg - rg_th)\n", " return {\"rg\":rg}\n", " self._callbacks[\"model\"][\"loss\"].append(loss_fn)\n", " self.opt[\"weights\"][\"rg\"] = weight\n", "\n", "def add_mpnn_loss(self, mpnn=0.1, mpnn_seq=0.0):\n", " '''\n", " add mpnn loss\n", " mpnn = maximize confidence of proteinmpnn\n", " mpnn_seq = push designed sequence to match proteinmpnn logits\n", " '''\n", "\n", " self._mpnn = mk_mpnn_model(weights = \"soluble\" if use_solubleMPNN else \"original\")\n", " def loss_fn(inputs, outputs, aux, key):\n", "\n", " # get structure\n", " atom_idx = tuple(residue_constants.atom_order[k] for k in [\"N\",\"CA\",\"C\",\"O\"])\n", " I = {\"S\": inputs[\"aatype\"],\n", " \"residue_idx\": inputs[\"residue_index\"],\n", " \"chain_idx\": inputs[\"asym_id\"],\n", " \"X\": outputs[\"structure_module\"][\"final_atom_positions\"][:,atom_idx],\n", " \"mask\": outputs[\"structure_module\"][\"final_atom_mask\"][:,1],\n", " \"lengths\": self._lengths,\n", " \"key\": key}\n", "\n", " if \"offset\" in inputs:\n", " I[\"offset\"] = inputs[\"offset\"]\n", "\n", " # set autoregressive mask\n", " L = sum(self._lengths)\n", " if self.protocol == \"binder\":\n", " I[\"ar_mask\"] = 1 - np.eye(L)\n", " I[\"ar_mask\"][-self._len:,-self._len:] = 0\n", " else:\n", " I[\"ar_mask\"] = np.zeros((L,L))\n", "\n", " # get logits\n", " logits = self._mpnn._score(**I)[\"logits\"][:,:20]\n", " if self.protocol == \"binder\":\n", " logits = logits[-self._len:]\n", " else:\n", " logits = logits[:self._len]\n", " aux[\"mpnn_logits\"] = logits\n", "\n", " # compute loss\n", " log_q = jax.nn.log_softmax(logits)\n", " p = inputs[\"seq\"][\"hard\"]\n", " q = jax.nn.softmax(logits)\n", " losses = {}\n", " losses[\"mpnn\"] = -log_q.max(-1).mean()\n", " losses[\"mpnn_seq\"] = -(p * jax.lax.stop_gradient(log_q)).sum(-1).mean()\n", " return losses\n", "\n", " self._callbacks[\"model\"][\"loss\"].append(loss_fn)\n", " self.opt[\"weights\"][\"mpnn\"] = mpnn\n", " self.opt[\"weights\"][\"mpnn_seq\"] = mpnn_seq\n", "\n", "clear_mem()\n", "af_model = mk_afdesign_model(protocol=\"hallucination\")\n", "af_model.prep_inputs(length=LENGTH, copies=COPIES)\n", "\n", "# add extra losses\n", "if use_rg_loss: add_rg_loss(af_model)\n", "if use_mpnn_loss: add_mpnn_loss(af_model)\n", "\n", "print(\"length\",af_model._lengths)\n", "print(\"weights\",af_model.opt[\"weights\"])" ] }, { "cell_type": "code", "source": [ "af_model.restart()\n", "if MODE == \"original\":\n", " # pre-design with gumbel initialization and softmax activation\n", " af_model.set_weights(plddt=0.0, pae=0.0)\n", " af_model.set_seq(mode=[\"gumbel\"])\n", " af_model.design_soft(50)\n", " af_model.set_seq(af_model.aux[\"seq\"][\"pseudo\"])\n", "\n", "if MODE == \"manuscript\":\n", " af_model.set_seq(mode=[\"gumbel\",\"soft\"])\n", "\n", "af_model.set_weights(plddt=1.0, pae=1.0)\n", "af_model.design_logits(40)\n", "af_model.design_logits(10, save_best=True)" ], "metadata": { "id": "f76xqCkw0vj9" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "A1GxeLZdTTya" }, "outputs": [], "source": [ "af_model.save_pdb(f\"{af_model.protocol}.pdb\")\n", "af_model.plot_pdb()" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "L2E9Tn2Acchj" }, "outputs": [], "source": [ "HTML(af_model.animate())" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "YSKWYu0_GlUH" }, "outputs": [], "source": [ "af_model.get_seqs()" ] }, { "cell_type": "code", "source": [ "#@markdown #Redesign with ProteinMPNN\n", "num_seqs = 8 #@param [\"8\", \"16\", \"32\", \"64\"] {type:\"raw\"}\n", "mpnn_sampling_temp = 0.1 #@param [\"0.0001\", \"0.1\", \"0.15\", \"0.2\", \"0.25\", \"0.3\", \"0.5\", \"1.0\"] {type:\"raw\"}\n", "rm_aa = \"C\" #@param {type:\"string\"}\n", "use_solubleMPNN = False #@param {type:\"boolean\"}\n", "#@markdown - `mpnn_sampling_temp` - control diversity of sampled sequences. (higher = more diverse).\n", "#@markdown - `rm_aa='C'` - do not use [C]ysteines.\n", "#@markdown - `use_solubleMPNN` - use weights trained only on soluble proteins. See [preprint](https://www.biorxiv.org/content/10.1101/2023.05.09.540044v2).\n", "#@markdown" ], "metadata": { "cellView": "form", "id": "m2qAYsDsCfqJ" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "from colabdesign.shared.protein import alphabet_list as chain_list\n", "mpnn_model = mk_mpnn_model()\n", "mpnn_model.prep_inputs(pdb_filename=f\"{af_model.protocol}.pdb\",\n", " chain=\",\".join(chain_list[:COPIES]),\n", " homooligmer=COPIES>1,\n", " rm_aa=rm_aa,\n", " weights = \"soluble\" if use_solubleMPNN else\"original\")\n", "out = mpnn_model.sample(num=num_seqs//8,\n", " batch=8,\n", " temperature=mpnn_sampling_temp)\n", "for seq,score in zip(out[\"seq\"],out[\"score\"]):\n", " print(score,seq.split(\"/\")[0])" ], "metadata": { "id": "uQa0FAp7bGQo" }, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "source": [ "#Run ESMfold" ], "metadata": { "id": "eDvyemgjNbX4" } }, { "cell_type": "code", "source": [ "print(\"# rmsd tmscore sequence\")\n", "best = {}\n", "best_rmsd = None\n", "for n,seq in enumerate(out[\"seq\"]):\n", " x = seq.split(\"/\")[0]\n", " with open(f\"{af_model.protocol}.esmfold.{n}.pdb\",\"w\") as handle:\n", " pdb_str = esmfold_api(x)\n", " handle.write(pdb_str)\n", " o = tmscore(f\"{af_model.protocol}.pdb\",\n", " f\"{af_model.protocol}.esmfold.{n}.pdb\")\n", " print(n,o[\"rms\"],o[\"tms\"],x)\n", " if best_rmsd is None or o[\"rms\"] < best_rmsd:\n", " best_rmsd = o[\"rms\"]\n", " best = {**o,\"seq\":x}" ], "metadata": { "id": "Ey29NmNAFtK0" }, "execution_count": null, "outputs": [] }, { "cell_type": "code", "source": [ "best" ], "metadata": { "id": "ltH6cLw5NhuX" }, "execution_count": null, "outputs": [] } ], "metadata": { "accelerator": "GPU", "colab": { "collapsed_sections": [ "q4qiU9I0QHSz" ], "provenance": [], "include_colab_link": true }, "kernelspec": { "display_name": "Python 3", "name": "python3" }, "language_info": { "name": "python" } }, "nbformat": 4, "nbformat_minor": 0 }