File size: 2,916 Bytes
918beed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
"""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
        )