File size: 4,180 Bytes
d53adc9
 
 
 
 
 
 
 
 
 
5ff48ac
d53adc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5ff48ac
 
 
 
d53adc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
# Exported for HuggingFace trust_remote_code loading.
# This file is intentionally self-contained.

"""HF adapter support for config metadata and remote-code registration."""

from __future__ import annotations

from collections.abc import Callable, Mapping
from dataclasses import asdict, is_dataclass
from typing import Protocol, TypeVar

from .canonical_config import SophiaModelConfig
from .hf_remote_code import (
    HF_CAUSAL_LM_CLASS,
    HF_CONFIG_CLASS,
    HF_MODELING_MODULE,
    HF_RUNTIME_FILENAME,
    HF_SUPPORT_FILENAME,
)


class AutoMapConfig(Protocol):
    architectures: list[str]
    auto_map: dict[str, str]


class ConfiguredModel(Protocol):
    config: AutoMapConfig | None


class MetadataConfig(AutoMapConfig, Protocol):
    max_seq_len: int
    dim: int
    n_layers: int
    num_heads: int
    norm_eps: float
    dropout: float
    head_dim: int
    max_position_embeddings: int
    hidden_size: int
    num_hidden_layers: int
    num_attention_heads: int
    sliding_window: int
    rms_norm_eps: float
    attention_dropout: float
    loss_chunk_size: int
    gradient_checkpointing_exclude_first: int
    gradient_checkpointing_exclude_last: int
    return_logits_in_train: bool
    use_cache: bool




def _to_dict_values(config: object) -> dict[str, object] | None:
    to_dict = getattr(config, "to_dict", None)
    if callable(to_dict):
        values = to_dict()
        if isinstance(values, Mapping):
            return dict(values)
    return None


def apply_auto_map(config: AutoMapConfig) -> AutoMapConfig:
    config.architectures = [HF_CAUSAL_LM_CLASS]
    config.auto_map = {
        "AutoConfig": f"{HF_MODELING_MODULE}.{HF_CONFIG_CLASS}",
        "AutoModelForCausalLM": f"{HF_MODELING_MODULE}.{HF_CAUSAL_LM_CLASS}",
    }
    return config


def ensure_auto_map(model: ConfiguredModel) -> None:
    cfg = model.config
    if cfg is None:
        raise RuntimeError("Model has no config; unable to set HuggingFace auto_map.")
    apply_auto_map(cfg)


def apply_config_metadata(
    config: MetadataConfig,
    *,
    return_logits_in_train: bool,
    use_cache: bool,
    loss_chunk_size: int,
    gradient_checkpointing_exclude_first: int,
    gradient_checkpointing_exclude_last: int,
) -> None:
    config.max_position_embeddings = int(config.max_seq_len)
    config.hidden_size = int(config.dim)
    config.num_hidden_layers = int(config.n_layers)
    config.num_attention_heads = int(config.num_heads)
    config.sliding_window = int(config.max_seq_len)
    config.rms_norm_eps = float(config.norm_eps)
    config.attention_dropout = float(config.dropout)
    config.loss_chunk_size = max(int(loss_chunk_size), 0)
    config.gradient_checkpointing_exclude_first = max(
        int(gradient_checkpointing_exclude_first),
        0,
    )
    config.gradient_checkpointing_exclude_last = max(
        int(gradient_checkpointing_exclude_last),
        0,
    )
    config.return_logits_in_train = bool(return_logits_in_train)
    config.use_cache = bool(use_cache)
    config.architectures = [HF_CAUSAL_LM_CLASS]


ConfigT = TypeVar("ConfigT")


def build_config(
    config: object,
    *,
    config_cls: Callable[..., ConfigT] | None = None,
) -> ConfigT:
    if config_cls is None:
        from .hf_projection import SophiaConfig

        config_cls = SophiaConfig

    if isinstance(config, Mapping):
        values = dict(config)
    elif is_dataclass(config):
        values = asdict(config)
    else:
        values = _to_dict_values(config)
        if values is None:
            canonical_fields = set(SophiaModelConfig.get_defaults())
            values = {
                key: getattr(config, key)
                for key in canonical_fields
                if hasattr(config, key)
            }

    return config_cls(**values)


__all__ = [
    "HF_CAUSAL_LM_CLASS",
    "HF_CONFIG_CLASS",
    "HF_MODELING_MODULE",
    "HF_RUNTIME_FILENAME",
    "HF_SUPPORT_FILENAME",
    "AutoMapConfig",
    "ConfiguredModel",
    "MetadataConfig",
    "apply_auto_map",
    "apply_config_metadata",
    "build_config",
    "ensure_auto_map",
]