Files
K3/kda/models/config.py
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

63 lines
2.1 KiB
Python

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