"""KDAConfig — toy Causal LM hyperparameters. Defaults match the working reference-backend model: GVA with G=2, safe gate (lower_bound=-5), q/k L2-norm and beta sigmoid inside the op. """ from __future__ import annotations from dataclasses import dataclass @dataclass class KDAConfig: hidden_size: int = 64 num_hidden_layers: int = 2 num_heads: int = 4 num_value_heads: int = 8 # G = num_value_heads // num_heads head_dim: int = 16 chunk_size: int = 16 vocab_size: int = 256 intermediate_size: int = 128 max_position_embeddings: int = 128 initializer_range: float = 0.02 norm_eps: float = 1e-6 use_gate_in_kernel: bool = True use_qk_l2norm_in_kernel: bool = True use_beta_sigmoid_in_kernel: bool = True lower_bound: float | None = -5.0 tie_word_embeddings: bool = False kda_backend: str = "reference" # reference | triton | fla attnres: str = "off" # off | full | block attnres_block_size: int | None = None # DecoderBlocks / block; None ≈ L/8 attnres_zero_init_queries: bool = True attnres_final_aggregate: bool = True gradient_checkpointing: bool = False @property def H(self) -> int: return self.num_heads @property def G(self) -> int: return self.num_value_heads // self.num_heads @property def HV(self) -> int: return self.num_value_heads @property def K(self) -> int: return self.head_dim @property def V(self) -> int: return self.head_dim def __post_init__(self): from ..layers.attn_res import validate_attnres if self.num_value_heads % self.num_heads: raise ValueError("num_value_heads must be divisible by num_heads") supported = {"reference", "triton", "fla", "torch", "auto"} if self.kda_backend not in supported: raise ValueError(f"kda_backend must be one of {sorted(supported)}") validate_attnres(self.attnres, self.attnres_block_size) def layer_specs(self) -> list[tuple[str, str]]: return [("kda", "swiglu")] * self.num_hidden_layers