ESMFold2-Fast / modeling_esmfold2_remote.py
fmilletari-czi's picture
Upload folder using huggingface_hub
918beed verified
Raw History Blame Contribute Delete
2.92 kB
"""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
)