{
"cells": [
{
"cell_type": "markdown",
"metadata": {
"id": "view-in-github",
"colab_type": "text"
},
"source": [
"
"
]
},
{
"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
}