OneScience's picture
Upload folder using huggingface_hub
35cdf53 verified
Raw
History Blame Contribute Delete
2.57 kB
"""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=[],
)