"""Load this checkpoint with the ``esm`` package instead of the transformers port.

Shipped inside each ``biohub/ESMC-*`` repo and named by its ``config.json``
``auto_map``, so ``AutoModel.from_pretrained(repo, trust_remote_code=True)``
returns an ``esm.models.esmc`` model. Without the flag transformers keeps using
its own port, which it has shipped since 5.x.

One shim class per Auto class transformers maps for ``esmc``. Dropping any of
them would break that Auto class on this repo: ``AutoConfig`` resolves to
:class:`EsmcRemoteConfig`, which is in none of transformers' own mappings.

Runs from the Hub module cache in the user's environment, so it imports nothing
but ``transformers``, ``packaging`` and ``esm``. The ESMFold2 loader is a
separate file for the same reason -- each repo downloads one module, so there is
nowhere to share code from.
"""

from typing import Any

from packaging.version import Version
from transformers.configuration_utils import PretrainedConfig
from transformers.modeling_utils import PreTrainedModel

#: Inclusive floor. 3.4.0 reads the published layout but takes no config
#: overrides, so it would silently ignore a ``num_labels`` on a head load.
MIN_ESM_VERSION = "3.4.1"

#: The keyword arguments esm's ESMC ``from_pretrained`` declares.
_ESM_KWARGS = frozenset(
    {
        "device",
        "dtype",
        "attn_implementation",
        "revision",
        "cache_dir",
        "token",
        "local_files_only",
        "force_download",
    }
)


def esmc_class(name: str) -> Any:
    """Return a class from ``esm.models.esmc``, or raise ``ImportError``.

    The message names the pip command; a bare ImportError out of the Hub module
    loader would not.
    """
    try:
        import esm
        from esm.models import esmc
    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 ESMC 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 getattr(esmc, name)


class EsmcRemoteConfig(PretrainedConfig):
    """Holds config.json verbatim. esm re-reads the file and builds its own config."""

    model_type = "esmc"


class _EsmcRemoteLoader(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 = EsmcRemoteConfig

    #: The ``esm.models.esmc`` class this shim loads.
    esm_class_name = "EsmcModel"

    #: esm keyword -> config attribute. transformers absorbs these into
    #: ``config`` before the model sees them, so
    #: ``AutoModel(..., attn_implementation="flash_attention_2")`` never
    #: reaches us as a keyword. Read back off that config instead of dropped.
    esm_config_fields: dict[str, str] = {"attn_implementation": "_attn_implementation"}

    @classmethod
    def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):  # type: ignore[override]
        # esm's signature is fixed, so transformers' trust_remote_code /
        # _from_auto / config would TypeError.
        forwarded = {k: v for k, v in kwargs.items() if k in _ESM_KWARGS}
        config = kwargs.get("config")
        for keyword, attribute in cls.esm_config_fields.items():
            value = getattr(config, attribute, None)
            if value is not None:
                forwarded.setdefault(keyword, value)
        return esmc_class(cls.esm_class_name).from_pretrained(
            pretrained_model_name_or_path, *args, **forwarded
        )


class EsmcRemoteModel(_EsmcRemoteLoader):
    """``AutoModel`` -- the bare encoder."""

    esm_class_name = "EsmcModel"


class EsmcRemoteForMaskedLM(_EsmcRemoteLoader):
    """``AutoModelForMaskedLM`` -- the encoder plus the published LM head."""

    esm_class_name = "EsmcForMaskedLM"


#: A head load also has to carry the shape of the head, which transformers
#: takes as ``num_labels`` and resolves against ``id2label``.
_HEAD_FIELDS = {
    **_EsmcRemoteLoader.esm_config_fields,
    "num_labels": "num_labels",
    "problem_type": "problem_type",
    "classifier_dropout": "classifier_dropout",
}


class EsmcRemoteForSequenceClassification(_EsmcRemoteLoader):
    """``AutoModelForSequenceClassification`` -- head is randomly initialised."""

    esm_class_name = "EsmcForSequenceClassification"
    esm_config_fields = _HEAD_FIELDS


class EsmcRemoteForTokenClassification(_EsmcRemoteLoader):
    """``AutoModelForTokenClassification`` -- head is randomly initialised."""

    esm_class_name = "EsmcForTokenClassification"
    esm_config_fields = _HEAD_FIELDS
