Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 41 additions & 12 deletions openfold/data/data_transforms_multimer.py
Original file line number Diff line number Diff line change
Expand Up @@ -318,19 +318,48 @@ def randint(lower, upper, generator, device):


def get_interface_residues(positions, atom_mask, asym_id, interface_threshold):
coord_diff = positions[..., None, :, :] - positions[..., None, :, :, :]
pairwise_dists = torch.sqrt(torch.sum(coord_diff ** 2, dim=-1))

diff_chain_mask = (asym_id[..., None, :] != asym_id[..., :, None]).float()
pair_mask = atom_mask[..., None, :] * atom_mask[..., None, :, :]
mask = (diff_chain_mask[..., None] * pair_mask).bool()

min_dist_per_res, _ = torch.where(mask, pairwise_dists, torch.inf).min(dim=-1)

valid_interfaces = torch.sum((min_dist_per_res < interface_threshold).float(), dim=-1)
interface_residues_idxs = torch.nonzero(valid_interfaces, as_tuple=True)[0]
num_res = positions.shape[0]
chunk_size = 64
atom_mask = atom_mask.bool()
interface_mask = torch.zeros(
num_res,
dtype=torch.bool,
device=positions.device,
)

return interface_residues_idxs
for query_start in range(0, num_res, chunk_size):
query_end = min(query_start + chunk_size, num_res)
query_positions = positions[query_start:query_end]
query_atom_mask = atom_mask[query_start:query_end]
query_asym_id = asym_id[query_start:query_end]

for key_start in range(0, num_res, chunk_size):
key_end = min(key_start + chunk_size, num_res)
key_positions = positions[key_start:key_end]
key_atom_mask = atom_mask[key_start:key_end]
key_asym_id = asym_id[key_start:key_end]

pairwise_dists = torch.cdist(
query_positions[:, None, :, :],
key_positions[None, :, :, :],
)
diff_chain_mask = (
query_asym_id[:, None] != key_asym_id[None, :]
)[:, :, None, None]
pair_mask = (
query_atom_mask[:, None, :, None]
& key_atom_mask[None, :, None, :]
)
contacts = (
(pairwise_dists < interface_threshold)
& diff_chain_mask
& pair_mask
)
interface_mask[query_start:query_end] |= contacts.any(
dim=(1, 2, 3)
)

return torch.nonzero(interface_mask, as_tuple=True)[0]


def get_spatial_crop_idx(protein, crop_size, interface_threshold, generator):
Expand Down
35 changes: 35 additions & 0 deletions tests/test_data_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
correct_msa_restypes, squeeze_features, randomly_replace_msa_with_unknown, MSA_FEATURE_NAMES, sample_msa, \
crop_extra_msa, delete_extra_msa, nearest_neighbor_clusters, make_msa_mask, make_hhblits_profile, make_masked_msa, \
make_msa_feat, crop_templates, make_atom14_masks
from openfold.data.data_transforms_multimer import get_interface_residues
from openfold.np import residue_constants as rc
from tests.config import config


Expand Down Expand Up @@ -221,6 +223,39 @@ def test_make_atom14_masks(self):
assert 'residx_atom37_to_atom14' in protein
assert 'atom37_atom_exists' in protein

def test_get_interface_residues(self):
positions = torch.zeros((2, rc.atom_type_num, 3))
atom_mask = torch.zeros((2, rc.atom_type_num))
asym_id = torch.tensor([0, 1])

lys_nz_idx = rc.atom_order["NZ"]
asp_od1_idx = rc.atom_order["OD1"]
positions[0, lys_nz_idx] = torch.tensor([0., 0., 0.])
positions[1, asp_od1_idx] = torch.tensor([1., 0., 0.])
atom_mask[0, lys_nz_idx] = 1.
atom_mask[1, asp_od1_idx] = 1.

interface_residues = get_interface_residues(
positions,
atom_mask,
asym_id,
interface_threshold=2.,
)

self.assertTrue(
torch.equal(interface_residues, torch.tensor([0, 1]))
)

positions[1, asp_od1_idx] = torch.tensor([3., 0., 0.])
interface_residues = get_interface_residues(
positions,
atom_mask,
asym_id,
interface_threshold=2.,
)

self.assertEqual(interface_residues.numel(), 0)


if __name__ == '__main__':
unittest.main()