Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
63 lines
2.1 KiB
Python
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
|