File size: 2,566 Bytes
35cdf53
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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


"""Batch dataclass."""
import dataclasses
from typing import Self

from flax_model.alphafold3.model import features
#import chex
import jax


@dataclasses.dataclass(frozen=True)
class Batch:
  """Dataclass containing batch."""

  msa: features.MSA
  templates: features.Templates
  token_features: features.TokenFeatures
  ref_structure: features.RefStructure
  predicted_structure_info: features.PredictedStructureInfo
  polymer_ligand_bond_info: features.PolymerLigandBondInfo
  ligand_ligand_bond_info: features.LigandLigandBondInfo
  pseudo_beta_info: features.PseudoBetaInfo
  atom_cross_att: features.AtomCrossAtt
  convert_model_output: features.ConvertModelOutput
  frames: features.Frames

  @property
  def num_res(self) -> int:
    return self.token_features.aatype.shape[-1]

  @classmethod
  def from_data_dict(cls, batch: features.BatchDict) -> Self:
    """Construct batch object from dictionary."""
    return cls(
        msa=features.MSA.from_data_dict(batch),
        templates=features.Templates.from_data_dict(batch),
        token_features=features.TokenFeatures.from_data_dict(batch),
        ref_structure=features.RefStructure.from_data_dict(batch),
        predicted_structure_info=features.PredictedStructureInfo.from_data_dict(
            batch
        ),
        polymer_ligand_bond_info=features.PolymerLigandBondInfo.from_data_dict(
            batch
        ),
        ligand_ligand_bond_info=features.LigandLigandBondInfo.from_data_dict(
            batch
        ),
        pseudo_beta_info=features.PseudoBetaInfo.from_data_dict(batch),
        atom_cross_att=features.AtomCrossAtt.from_data_dict(batch),
        convert_model_output=features.ConvertModelOutput.from_data_dict(batch),
        frames=features.Frames.from_data_dict(batch),
    )

  def as_data_dict(self) -> features.BatchDict:
    """Converts batch object to dictionary."""
    output = {
        **self.msa.as_data_dict(),
        **self.templates.as_data_dict(),
        **self.token_features.as_data_dict(),
        **self.ref_structure.as_data_dict(),
        **self.predicted_structure_info.as_data_dict(),
        **self.polymer_ligand_bond_info.as_data_dict(),
        **self.ligand_ligand_bond_info.as_data_dict(),
        **self.pseudo_beta_info.as_data_dict(),
        **self.atom_cross_att.as_data_dict(),
        **self.convert_model_output.as_data_dict(),
        **self.frames.as_data_dict(),
    }
    return output


jax.tree_util.register_dataclass(
    Batch,
    data_fields=[f.name for f in dataclasses.fields(Batch)],
    meta_fields=[],
)