kda-sigmoid-hybrid-1.3B-100B / configuration_complex_kda.py
korbip's picture
Add ComplexKDA 1.3B/100BT FineWeb-Edu model
92db682 verified
Raw History Blame
10 kB
# 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)