File size: 10,005 Bytes
92db682
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
# Copyright (c) 2026, the ComplexKDA authors.
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li (the parts derived
# from flash-linear-attention, MIT licensed).
#
# SPDX-License-Identifier: MIT
"""Config for the ComplexKDA models published on the Hub.

STANDALONE ON PURPOSE. `fla` is not imported anywhere in this file, so the
config loads with nothing but `transformers` installed. The hybrid-attention
normalisation that fla keeps in `fla/models/hybrid.py` is vendored below for
the same reason -- a config that needed the fork to parse would make every
error message about a missing package rather than about the model.
"""

from __future__ import annotations

import math

from transformers.configuration_utils import PretrainedConfig

__all__ = ["ComplexKDAConfig"]


# ---------------------------------------------------------------------------
# hybrid attention spec (vendored from fla/models/hybrid.py)
#
# A hybrid arm replaces the linear mixer with full attention at the layers
# named in `attn["layers"]`. The dict is stored in config.json, so it is
# validated on the way in: an out-of-range or duplicated layer index would
# otherwise build a model whose `state_dict` silently disagrees with the
# checkpoint at exactly those layers.
# ---------------------------------------------------------------------------

def _spec_context(spec_index: int | None) -> str:
    return "attn specification" if spec_index is None else f"attn specification at index {spec_index}"


def _positive_int(value: object, *, field: str, context: str) -> int:
    if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
        raise ValueError(f"{context} field {field!r} must be a positive integer; got {value!r}")
    return value


def _normalize_spec(spec: dict, *, num_hidden_layers: int, spec_index: int | None,
                    assigned_layers: dict) -> dict:
    context = _spec_context(spec_index)
    normalized = dict(spec)

    for field in ("layers", "num_heads"):
        if field not in normalized:
            raise ValueError(f"{context} field {field!r} is required; got <missing>")

    layers = normalized["layers"]
    if not isinstance(layers, (list, tuple)):
        raise ValueError(
            f"{context} field 'layers' must be a list or tuple of integer layer indices; got {layers!r}")

    normalized_layers, seen = [], set()
    for layer_idx in layers:
        if isinstance(layer_idx, bool) or not isinstance(layer_idx, int):
            raise ValueError(f"{context} field 'layers' must contain only integer layer indices; got {layer_idx!r}")
        if layer_idx < 0 or layer_idx >= num_hidden_layers:
            raise ValueError(
                f"{context} field 'layers' contains out-of-range layer {layer_idx!r}; "
                f"expected a value in [0, {num_hidden_layers})")
        if layer_idx in seen:
            raise ValueError(f"{context} field 'layers' contains duplicate layer {layer_idx!r}; got {layers!r}")
        if layer_idx in assigned_layers:
            raise ValueError(
                f"{context} assigns conflicting layer {layer_idx!r}, which is already assigned by "
                f"{_spec_context(assigned_layers[layer_idx])}")
        seen.add(layer_idx)
        assigned_layers[layer_idx] = spec_index
        normalized_layers.append(layer_idx)

    normalized["layers"] = normalized_layers
    normalized["num_heads"] = _positive_int(normalized["num_heads"], field="num_heads", context=context)

    num_kv_heads = normalized.get("num_kv_heads")
    if num_kv_heads is None:
        num_kv_heads = normalized["num_heads"]
    normalized["num_kv_heads"] = _positive_int(num_kv_heads, field="num_kv_heads", context=context)

    qkv_bias = normalized.get("qkv_bias", False)
    if not isinstance(qkv_bias, bool):
        raise ValueError(f"{context} field 'qkv_bias' must be a Boolean; got {qkv_bias!r}")
    normalized["qkv_bias"] = qkv_bias

    window_size = normalized.get("window_size")
    if window_size is not None:
        window_size = _positive_int(window_size, field="window_size", context=context)
    normalized["window_size"] = window_size

    rope_theta = normalized.get("rope_theta", 10000.0)
    try:
        ok = (not isinstance(rope_theta, bool) and isinstance(rope_theta, (int, float))
              and math.isfinite(rope_theta) and rope_theta > 0)
    except OverflowError:
        ok = False
    if not ok:
        raise ValueError(f"{context} field 'rope_theta' must be positive and finite; got {rope_theta!r}")
    normalized["rope_theta"] = rope_theta

    return normalized


def normalize_hybrid_attention_config(attn, *, num_hidden_layers: int):
    """Validate and normalise `attn`: None, one spec dict, or a list of them."""
    if attn is None:
        return None
    if isinstance(num_hidden_layers, bool) or not isinstance(num_hidden_layers, int) or num_hidden_layers < 0:
        raise ValueError(f"field 'num_hidden_layers' must be a non-negative integer; got {num_hidden_layers!r}")
    if not isinstance(attn, (dict, list)):
        raise ValueError(f"attn must be None, a dictionary, or a list of dictionaries; got {attn!r}")

    is_single = isinstance(attn, dict)
    specs = [attn] if is_single else attn
    assigned: dict = {}
    out = []
    for i, spec in enumerate(specs):
        idx = None if is_single else i
        if not isinstance(spec, dict):
            raise ValueError(f"{_spec_context(idx)} must be a dictionary; got {spec!r}")
        out.append(_normalize_spec(spec, num_hidden_layers=num_hidden_layers,
                                   spec_index=idx, assigned_layers=assigned))
    return out[0] if is_single else out


def get_hybrid_attention_spec(attn, *, layer_idx: int):
    """The normalised spec assigned to `layer_idx`, or None for a linear layer."""
    if attn is None:
        return None
    for spec in ([attn] if isinstance(attn, dict) else attn):
        if layer_idx in spec["layers"]:
            return spec
    return None


class ComplexKDAConfig(PretrainedConfig):
    """Kimi Delta Attention with a *signed* (complex, i.e. Z_2-phased) decay gate.

    The baseline arms set ``gate="sigmoid"`` with ``allow_neg_eigval=False``;
    the ComplexKDA arms set ``gate="signed_sigmoid2"`` with
    ``allow_neg_eigval=True``. Everything else is shared, which is what makes
    the two comparable.
    """

    model_type = "complex_kda"
    keys_to_ignore_at_inference = ["past_key_values"]

    def __init__(
        self,
        attn_mode: str = "chunk",
        hidden_size: int = 2048,
        expand_v: float = 1.0,
        use_short_conv: bool = True,
        drop_silu: bool = False,
        drop_key_silu: bool = False,
        conv_silu: str = "qkv",
        allow_neg_eigval: bool = True,
        gate: str = "signed_sigmoid2",
        gate_init_style: str = "shipped",
        output_gate: str = "lowrank",
        beta_init_style: str = "standard",
        lower_bound: float = -5.0,
        num_heads: int = 16,
        num_v_heads: int | None = None,
        head_dim: int = 128,
        num_hidden_layers: int = 24,
        norm_eps: float = 1e-6,
        conv_size: int = 4,
        attn: dict | list | None = None,
        hidden_ratio: int | None = 4,
        intermediate_size: int | None = None,
        hidden_act: str = "swish",
        max_position_embeddings: int = 4096,
        initializer_range: float = 0.02,
        vocab_size: int = 32000,
        tie_word_embeddings: bool = False,
        use_cache: bool = True,
        pad_token_id: int | None = None,
        bos_token_id: int = 1,
        eos_token_id: int = 2,
        # Kept so a config written by the training stack round-trips. They
        # select fused kernels when the fla fork is installed and are ignored
        # by the pure-torch path, which computes the same thing either way.
        fuse_norm: bool = True,
        fuse_swiglu: bool = True,
        fuse_cross_entropy: bool = True,
        use_l2warp: bool = False,
        **kwargs,
    ):
        self.attn_mode = attn_mode
        self.hidden_size = hidden_size
        self.expand_v = expand_v
        self.use_short_conv = use_short_conv
        self.drop_silu = drop_silu
        self.drop_key_silu = drop_key_silu
        self.conv_silu = conv_silu
        self.allow_neg_eigval = allow_neg_eigval
        self.gate = gate
        self.gate_init_style = gate_init_style
        self.output_gate = output_gate
        self.beta_init_style = beta_init_style
        self.lower_bound = lower_bound
        self.num_heads = num_heads
        self.num_v_heads = num_v_heads
        self.head_dim = head_dim
        # `num_hidden_layers` must be set before `attn`: the layer-range check
        # in the setter below reads it.
        self.num_hidden_layers = num_hidden_layers
        self.norm_eps = norm_eps
        self.conv_size = conv_size
        self.attn = attn

        self.hidden_ratio = hidden_ratio
        self.intermediate_size = intermediate_size
        self.hidden_act = hidden_act
        self.max_position_embeddings = max_position_embeddings
        self.initializer_range = initializer_range
        self.vocab_size = vocab_size
        self.use_cache = use_cache

        self.fuse_norm = fuse_norm
        self.fuse_swiglu = fuse_swiglu
        self.fuse_cross_entropy = fuse_cross_entropy
        self.use_l2warp = use_l2warp

        super().__init__(
            pad_token_id=pad_token_id,
            bos_token_id=bos_token_id,
            eos_token_id=eos_token_id,
            tie_word_embeddings=tie_word_embeddings,
            **kwargs,
        )

    # `attn` is a property so that assigning it after construction -- which
    # `PretrainedConfig.from_dict` does -- is validated too.
    @property
    def attn(self):
        return self.__dict__.get("attn")

    @attn.setter
    def attn(self, value) -> None:
        self.__dict__["attn"] = normalize_hybrid_attention_config(
            value, num_hidden_layers=self.num_hidden_layers)