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
        )