MilaDeepGraph/ProtST-ESM1b
025
1from transformers import PretrainedConfig2from transformers.utils import logging3from transformers.models.esm import EsmConfig4from transformers.models.bert import BertConfig5 6logger = logging.get_logger(__name__)7 8 9class ProtSTConfig(PretrainedConfig):10 r"""11 This is the configuration class to store the configuration of a [`ProtSTModel`].12 13 Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the14 documentation from [`PretrainedConfig`] for more information.15 16 Args:17 protein_config (`dict`, *optional*):18 Dictionary of configuration options used to initialize [`EsmForProteinRepresentation`].19 text_config (`dict`, *optional*):20 Dictionary of configuration options used to initialize [`BertForPubMed`].21 ```"""22 23 model_type = "protst"24 25 def __init__(26 self,27 protein_config=None,28 text_config=None,29 **kwargs,30 ):31 super().__init__(**kwargs)32 33 if protein_config is None:34 protein_config = {}35 logger.info("`protein_config` is `None`. Initializing the `ProtSTTextConfig` with default values.")36 37 if text_config is None:38 text_config = {}39 logger.info("`text_config` is `None`. Initializing the `ProtSTVisionConfig` with default values.")40 41 self.protein_config = EsmConfig(**protein_config)42 self.text_config = BertConfig(**text_config)43 44 @classmethod45 def from_protein_text_configs(46 cls, protein_config: EsmConfig, text_config: BertConfig, **kwargs47 ):48 r"""49 Instantiate a [`ProtSTConfig`] (or a derived class) from ProtST text model configuration. Returns:50 [`ProtSTConfig`]: An instance of a configuration object51 """52 53 return cls(protein_config=protein_config.to_dict(), text_config=text_config.to_dict(), **kwargs)