CoolFace
Modelpublic

MilaDeepGraph/ProtST-ESM1b

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes25downloads
configuration_protst.py53 linesDownload Raw Back to root
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)