"""Load this checkpoint with the ``esm`` package instead of the transformers port. Shipped inside each ``biohub/ESMFold2*`` repo and named by its ``config.json`` ``auto_map``, so ``AutoModel.from_pretrained(repo, trust_remote_code=True)`` returns an ``esm.models.esmfold2`` model. Without the flag transformers keeps using its own port. Runs from the Hub module cache in the user's environment, so it imports nothing but ``transformers``, ``packaging`` and ``esm``. """ from typing import Any from packaging.version import Version from transformers.configuration_utils import PretrainedConfig from transformers.modeling_utils import PreTrainedModel #: Inclusive floor. Older esm cannot read the bundled single-checkpoint layout. MIN_ESM_VERSION = "3.4.1" #: Everything esm's from_pretrained consumes, directly or via resolve_model_dir. _ESM_KWARGS = frozenset( { "load_esmc", "esmc_precision", "device", "dtype", "revision", "cache_dir", "token", "local_files_only", "force_download", } ) def esmfold2_class() -> Any: """Return esm's ``EsmFold2Model``, or raise ``ImportError`` naming the pip command.""" try: import esm from esm.models.esmfold2 import EsmFold2Model except ImportError as exc: raise ImportError( f"trust_remote_code=True runs the esm package, which failed to " f"import ({exc}). Install it with:\n\n" f" pip install 'esm>={MIN_ESM_VERSION}'\n\n" "Or drop the flag to use the ESMFold2 port in transformers." ) from exc if Version(esm.__version__) < Version(MIN_ESM_VERSION): raise ImportError( f"esm {esm.__version__} is installed, but this checkpoint needs " f"{MIN_ESM_VERSION} or newer. Upgrade with:\n\n" f" pip install --upgrade 'esm>={MIN_ESM_VERSION}'" ) return EsmFold2Model class EsmFold2RemoteConfig(PretrainedConfig): """Holds config.json verbatim. esm re-reads the file and builds its own config.""" model_type = "esmfold2" class EsmFold2RemoteModel(PreTrainedModel): """Loader shim; ``from_pretrained`` returns an esm model, not one of these. A ``PreTrainedModel`` subclass only because the Auto classes call ``register_for_auto_class`` on what they load and check its ``config_class``. """ config_class = EsmFold2RemoteConfig @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs): # type: ignore[override] # esm forwards leftover kwargs to resolve_model_dir, whose signature is # fixed, so transformers' trust_remote_code / _from_auto would TypeError. forwarded = {k: v for k, v in kwargs.items() if k in _ESM_KWARGS} return esmfold2_class().from_pretrained( pretrained_model_name_or_path, *args, **forwarded )