File size: 7,693 Bytes
fae1173 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 | import biotite.structure
import numpy as np
import torch
from typing import Sequence, Tuple, List
from util import (
load_structure,
extract_coords_from_structure,
load_coords,
get_sequence_loss,
get_encoder_output,
)
import util
def extract_coords_from_complex(structure: biotite.structure.AtomArray):
"""
Args:
structure: biotite AtomArray
Returns:
Tuple (coords_list, seq_list)
- coords: Dictionary mapping chain ids to L x 3 x 3 array for N, CA, C
coordinates representing the backbone of each chain
- seqs: Dictionary mapping chain ids to native sequences of each chain
"""
coords = {}
seqs = {}
all_chains = biotite.structure.get_chains(structure)
for chain_id in all_chains:
chain = structure[structure.chain_id == chain_id]
coords[chain_id], seqs[chain_id] = extract_coords_from_structure(chain)
return coords, seqs
def load_complex_coords(fpath, chains):
"""
Args:
fpath: filepath to either pdb or cif file
chains: the chain ids (the order matters for autoregressive model)
Returns:
Tuple (coords_list, seq_list)
- coords: Dictionary mapping chain ids to L x 3 x 3 array for N, CA, C
coordinates representing the backbone of each chain
- seqs: Dictionary mapping chain ids to native sequences of each chain
"""
structure = load_structure(fpath, chains)
return extract_coords_from_complex(structure)
#*
def _concatenate_coords(
coords,
target_chain_id,
padding_length=10,
order=None
):
"""
Args:
coords: Dictionary mapping chain ids to L x 3 x 3 array for N, CA, C
coordinates representing the backbone of each chain
target_chain_id: The chain id to sample sequences for
padding_length: Length of padding between concatenated chains
Returns:
Tuple (coords, seq)
- coords_concatenated is an L x 3 x 3 array for N, CA, C coordinates, a
concatenation of the chains with padding in between
AND target chain placed first
- seq is the extracted sequence, with padding tokens inserted
between the concatenated chains
"""
pad_coords = np.full((padding_length, 3, 3), np.nan, dtype=np.float32)
if order is None:
order = (
[ target_chain_id ] +
[ chain_id for chain_id in coords if chain_id != target_chain_id ]
)
coords_list, coords_chains = [], []
for idx, chain_id in enumerate(order):
if idx > 0:
coords_list.append(pad_coords)
coords_chains.append([ 'pad' ] * padding_length)
coords_list.append(list(coords[chain_id]))
coords_chains.append([ chain_id ] * coords[chain_id].shape[0])
coords_concatenated = np.concatenate(coords_list, axis=0)
coords_chains = np.concatenate(coords_chains, axis=0).ravel()
return coords_concatenated, coords_chains
#*
def _concatenate_seqs(
native_seqs,
target_seq,
target_chain_id,
padding_length=10,
order=None,
):
"""
Args:
native_seqs: Dictionary mapping chain ids to corresponding AA sequence
target_seq: The chain id to sample sequences for
padding_length: Length of padding between concatenated chains
Returns:
native_seqs_concatenated: Array of length L, concatenation of the chain
sequences with padding in between
"""
if order is None:
order = (
[ target_chain_id ] +
[ chain_id for chain_id in native_seqs if chain_id != target_chain_id ]
)
native_seqs_list = []
for idx, chain_id in enumerate(order):
if idx > 0:
native_seqs_list.append(['<mask>'] * (padding_length - 1) + ['<cath>'])
if chain_id == target_chain_id:
native_seqs_list.append(list(target_seq))
else:
native_seqs_list.append(list(native_seqs[chain_id]))
native_seqs_concatenated = ''.join(np.concatenate(native_seqs_list, axis=0))
return native_seqs_concatenated
#*
def sample_sequence_in_complex(model, coords, target_chain_id, temperature=1.,
padding_length=10):
"""
Samples sequence for one chain in a complex.
Args:
model: An instance of the GVPTransformer model
coords: Dictionary mapping chain ids to L x 3 x 3 array for N, CA, C
coordinates representing the backbone of each chain
target_chain_id: The chain id to sample sequences for
padding_length: padding length in between chains
Returns:
Sampled sequence for the target chain
"""
target_chain_len = coords[target_chain_id].shape[0]
all_coords, coords_chains = _concatenate_coords(coords, target_chain_id)
device = next(model.parameters()).device
# Supply padding tokens for other chains to avoid unused sampling for speed
padding_pattern = ['<pad>'] * all_coords.shape[0]
for i in range(target_chain_len):
padding_pattern[i] = '<mask>'
sampled = model.sample(all_coords, partial_seq=padding_pattern,
temperature=temperature, device=device)
sampled = sampled[:target_chain_len]
return sampled
#*
def score_sequence_in_complex(
model,
alphabet,
coords,
native_seqs,
target_chain_id,
target_seq,
padding_length=10,
order=None,
):
"""
Scores sequence for one chain in a complex.
Args:
model: An instance of the GVPTransformer model
alphabet: Alphabet for the model
coords: Dictionary mapping chain ids to L x 3 x 3 array for N, CA, C
coordinates representing the backbone of each chain
native_seqs: Dictionary mapping chain ids to sequence
extracted from each chain
target_chain_id: The chain id to sample sequences for
target_seq: Target sequence for the target chain for scoring.
padding_length: padding length in between chains
Returns:
Tuple (ll_fullseq, ll_withcoord)
- ll_fullseq: Average log-likelihood over the full target chain
- ll_targetseq Average log-likelihood in target chain excluding those
residues without coordinates
"""
assert(len(target_seq) == len(native_seqs[target_chain_id]))
all_coords, coords_chains = _concatenate_coords(
coords,
target_chain_id,
order=order,
)
all_seqs = _concatenate_seqs(
native_seqs,
target_seq,
target_chain_id,
order=order,
)
loss, target_padding_mask = get_sequence_loss(model, alphabet, all_coords,
all_seqs)
assert(all_coords.shape[0] == coords_chains.shape[0] == loss.shape[0])
ll_fullseq = -np.mean(loss[coords_chains != 'pad'])
ll_targetseq = -np.mean(loss[coords_chains == target_chain_id])
return ll_fullseq, ll_targetseq
def get_encoder_output_for_complex(model, alphabet, coords, target_chain_id):
"""
Args:
model: An instance of the GVPTransformer model
alphabet: Alphabet for the model
coords: Dictionary mapping chain ids to L x 3 x 3 array for N, CA, C
coordinates representing the backbone of each chain
target_chain_id: The chain id to sample sequences for
Returns:
Dictionary mapping chain id to encoder output for each chain
"""
all_coords = _concatenate_coords(coords, target_chain_id)
all_rep = get_encoder_output(model, alphabet, all_coords)
target_chain_len = coords[target_chain_id].shape[0]
return all_rep[:target_chain_len] |