# 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 ") 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)