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.
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user