File size: 6,384 Bytes
3e62986 | 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 |
"""Utilities for structure module testing."""
import dataclasses
from absl.testing import parameterized
from flax_model.alphafold3 import structure
from flax_model.alphafold3.common.testing import data
import numpy as np
import tree
class StructureTestCase(parameterized.TestCase):
"""Testing utilities for working with structure.Structure."""
def assertAuthorNamingSchemeEqual(self, ans1, ans2): # pylint: disable=invalid-name
"""Walks naming scheme, making sure all elements are equal."""
if ans1 is None or ans2 is None:
self.assertIsNone(ans1)
self.assertIsNone(ans2)
return
flat_ans1 = dict(tree.flatten_with_path(dataclasses.asdict(ans1)))
flat_ans2 = dict(tree.flatten_with_path(dataclasses.asdict(ans2)))
for k, v in flat_ans1.items():
self.assertEqual(v, flat_ans2[k], msg=str(k))
for k, v in flat_ans2.items():
self.assertEqual(v, flat_ans1[k], msg=str(k))
def assertAllResiduesEqual(self, all_res1, all_res2): # pylint: disable=invalid-name
"""Walks all residues, making sure alll elements are equal."""
if all_res1 is None or all_res2 is None:
self.assertIsNone(all_res1)
self.assertIsNone(all_res2)
return
self.assertSameElements(all_res1.keys(), all_res2.keys())
for chain_id, chain_res in all_res1.items():
self.assertSequenceEqual(chain_res, all_res2[chain_id], msg=chain_id)
def assertBioassemblyDataEqual(self, data1, data2): # pylint: disable=invalid-name
if data1 is None or data2 is None:
self.assertIsNone(data1)
self.assertIsNone(data2)
return
self.assertDictEqual(data1.to_mmcif_dict(), data2.to_mmcif_dict())
def assertChemicalComponentsDataEqual( # pylint: disable=invalid-name
self,
data1,
data2,
allow_chem_comp_data_extension,
):
"""Checks whether two ChemicalComponentData objects are considered equal."""
if data1 is None or data2 is None:
self.assertIsNone(data1)
self.assertIsNone(data2)
return
if (not allow_chem_comp_data_extension) or (
data1.chem_comp.keys() ^ data2.chem_comp.keys()
):
self.assertDictEqual(data1.chem_comp, data2.chem_comp)
else:
mismatching_values = []
for component_id in data1.chem_comp:
found = data1.chem_comp[component_id]
expected = data2.chem_comp[component_id]
if not found.extends(expected):
mismatching_values.append((component_id, expected, found))
if mismatching_values:
mismatch_err_msgs = '\n'.join(
f'{component_id}: {expected} or its extension expected,'
f' but {found} found.'
for component_id, expected, found in mismatching_values
)
self.fail(
f'Mismatching values for `_chem_comp` table: {mismatch_err_msgs}',
)
def assertBondsEqual(self, bonds1, bonds2, atom_key1, atom_key2): # pylint: disable=invalid-name
"""Checks whether two Bonds objects are considered equal."""
# An empty bonds table is functionally equivalent to an empty bonds table.
# NB: this can only ever be None in structure v1.
if bonds1 is None or not bonds1.size or bonds2 is None or not bonds2.size:
self.assertTrue(bonds1 is None or not bonds1.size, msg=f'{bonds1=}')
self.assertTrue(bonds2 is None or not bonds2.size, msg=f'{bonds2=}')
return
ptnr1_indices1, ptnr2_indices1 = bonds1.get_atom_indices(atom_key1)
ptnr1_indices2, ptnr2_indices2 = bonds2.get_atom_indices(atom_key2)
np.testing.assert_array_equal(ptnr1_indices1, ptnr1_indices2)
np.testing.assert_array_equal(ptnr2_indices1, ptnr2_indices2)
np.testing.assert_array_equal(bonds1.type, bonds2.type)
np.testing.assert_array_equal(bonds1.role, bonds2.role)
def assertStructuresEqual( # pylint: disable=invalid-name
self,
struc1,
struc2,
*,
ignore_fields=None,
allow_chem_comp_data_extension=False,
atol=0,
):
"""Checks whether two Structure objects could be considered equal.
Args:
struc1: First Structure object.
struc2: Second Structure object.
ignore_fields: Fields not taken into account during comparison.
allow_chem_comp_data_extension: Whether to allow data of `_chem_comp`
table to differ if `struc2` is missing some fields, but `struc1` has
specific values for them.
atol: Absolute tolerance for floating point comparisons (in
np.testing.assert_allclose).
"""
for field in sorted(structure.GLOBAL_FIELDS):
if ignore_fields and field in ignore_fields:
continue
if field == 'author_naming_scheme':
self.assertAuthorNamingSchemeEqual(struc1[field], struc2[field])
elif field == 'all_residues':
self.assertAllResiduesEqual(struc1[field], struc2[field])
elif field == 'bioassembly_data':
self.assertBioassemblyDataEqual(struc1[field], struc2[field])
elif field == 'chemical_components_data':
self.assertChemicalComponentsDataEqual(
struc1[field], struc2[field], allow_chem_comp_data_extension
)
elif field == 'bonds':
self.assertBondsEqual(
struc1.bonds, struc2.bonds, struc1.atom_key, struc2.atom_key
)
else:
self.assertEqual(struc1[field], struc2[field], msg=field)
# The chain order within a structure is arbitrary so in order to
# directly compare arrays we first align struc1 to struc2 and check that
# the number of atoms doesn't change.
num_atoms = struc1.num_atoms
self.assertEqual(struc2.num_atoms, num_atoms)
struc1 = struc1.order_and_drop_atoms_to_match(struc2)
self.assertEqual(struc1.num_atoms, num_atoms)
for field in sorted(structure.ARRAY_FIELDS):
if field == 'atom_key':
# atom_key has no external meaning, so it doesn't matter whether it
# differs between two structures.
continue
if ignore_fields and field in ignore_fields:
continue
self.assertEqual(struc1[field] is None, struc2[field] is None, msg=field)
if np.issubdtype(struc1[field].dtype, np.inexact):
np.testing.assert_allclose(
struc1[field], struc2[field], err_msg=field, atol=atol
)
else:
np.testing.assert_array_equal(
struc1[field], struc2[field], err_msg=field
)
|