Haiku-base / tiny_gdn /config.py
kerzgrr's picture
Upload Haiku-base (pretrain EMA @ step 8400)
fba158e verified
Raw History Blame Contribute Delete
8.14 kB
from __future__ import annotations
import json
from dataclasses import asdict, dataclass, fields
from pathlib import Path
from typing import Any
@dataclass(frozen=True)
class TinyGDNConfig:
architecture: str = "TinyGDNForCausalLM"
model_type: str = "tiny_gdn"
vocab_size: int = 49_152
# Deep-thin sizing is deliberate: controlled sub-billion studies find
# depth materially more valuable than width around the 125M-150M scale.
hidden_size: int = 512
intermediate_size: int = 1_472
num_hidden_layers: int = 32
num_attention_heads: int = 4
num_key_value_heads: int = 1
attention_head_dim: int = 128
full_attention_interval: int = 4
attention_dropout: float = 0.0
partial_rotary_factor: float = 0.5
rope_theta: float = 1_000_000.0
linear_num_heads: int = 4
linear_num_value_heads: int = 4
linear_head_dim: int = 128
linear_expand_v: float = 1.0
linear_conv_kernel_dim: int = 4
allow_negative_eigenvalues: bool = False
linear_layer_type: str = "gdn2"
attention_layer_type: str = "full_attention"
mlp_activation: str = "swiglu"
use_block_attn_res: bool = False
kda_safe_gate: bool = True
kda_lower_bound: float = -5.0
mla_q_lora_rank: int | None = 512
mla_kv_lora_rank: int = 256
mla_qk_nope_head_dim: int = 128
mla_v_head_dim: int = 128
situ_gate_cap: float = 4.0
situ_up_cap: float = 25.0
max_position_embeddings: int = 32_768
training_sequence_length: int = 2_048
rms_norm_eps: float = 1e-6
initializer_range: float = 0.02
tie_word_embeddings: bool = True
shared_layer_indices: tuple[int, ...] = ()
# MTP is an opt-in ablation at this scale; static MTP is not assumed to
# improve a 150M model without a controlled pilot.
mtp_num_heads: int = 0
mtp_adapter_rank: int = 128
mtp_loss_weight: float = 0.0
bos_token_id: int = 0
eos_token_id: int = 1
pad_token_id: int = 2
unk_token_id: int = 3
def __post_init__(self) -> None:
if self.vocab_size <= 0 or self.vocab_size > 65_536:
raise ValueError("vocab_size must fit the uint16 token dataset")
if self.hidden_size != self.num_attention_heads * self.attention_head_dim:
raise ValueError("hidden_size must equal num_attention_heads * attention_head_dim")
if self.hidden_size != self.linear_num_heads * self.linear_head_dim:
raise ValueError("hidden_size must equal linear_num_heads * linear_head_dim")
if self.linear_num_value_heads < self.linear_num_heads:
raise ValueError("linear_num_value_heads must be at least linear_num_heads")
if self.linear_num_value_heads % self.linear_num_heads != 0:
raise ValueError("linear_num_value_heads must be divisible by linear_num_heads")
if self.num_attention_heads % self.num_key_value_heads != 0:
raise ValueError("num_attention_heads must be divisible by num_key_value_heads")
if self.num_hidden_layers % self.full_attention_interval != 0:
raise ValueError("num_hidden_layers must be divisible by full_attention_interval")
if not 0.0 < self.partial_rotary_factor <= 1.0:
raise ValueError("partial_rotary_factor must be in (0, 1]")
rotary_dim = int(self.attention_head_dim * self.partial_rotary_factor)
if rotary_dim <= 0 or rotary_dim % 2:
raise ValueError("The partial rotary dimension must be positive and even")
if self.training_sequence_length > self.max_position_embeddings:
raise ValueError("training_sequence_length exceeds max_position_embeddings")
if len(set(self.shared_layer_indices)) != len(self.shared_layer_indices):
raise ValueError("shared_layer_indices must be unique")
if any(
index < 0 or index >= self.num_hidden_layers
for index in self.shared_layer_indices
):
raise ValueError("shared_layer_indices contains an invalid layer")
if self.mtp_num_heads < 0:
raise ValueError("mtp_num_heads cannot be negative")
if self.mtp_num_heads and self.mtp_adapter_rank <= 0:
raise ValueError("mtp_adapter_rank must be positive when MTP is enabled")
if not 0.0 <= self.mtp_loss_weight <= 1.0:
raise ValueError("mtp_loss_weight must be between zero and one")
for token_id in (
self.bos_token_id,
self.eos_token_id,
self.pad_token_id,
self.unk_token_id,
):
if not 0 <= token_id < self.vocab_size:
raise ValueError(f"Special token ID {token_id} is outside the vocabulary")
if self.linear_layer_type not in {"gdn2", "kda"}:
raise ValueError("linear_layer_type must be 'gdn2' or 'kda'")
if self.attention_layer_type not in {"full_attention", "gated_mla"}:
raise ValueError("attention_layer_type must be 'full_attention' or 'gated_mla'")
if self.mlp_activation not in {"swiglu", "situ_glu"}:
raise ValueError("mlp_activation must be 'swiglu' or 'situ_glu'")
if self.mla_kv_lora_rank <= 0:
raise ValueError("mla_kv_lora_rank must be positive")
if self.mla_q_lora_rank is not None and self.mla_q_lora_rank <= 0:
raise ValueError("mla_q_lora_rank must be positive or None")
if self.mla_qk_nope_head_dim <= 0 or self.mla_v_head_dim <= 0:
raise ValueError("MLA head dimensions must be positive")
if self.kda_lower_bound >= 0.0:
raise ValueError("kda_lower_bound must be negative log-decay")
if self.situ_gate_cap <= 0.0 or self.situ_up_cap <= 0.0:
raise ValueError("SiTU caps must be positive")
@property
def layer_types(self) -> tuple[str, ...]:
return tuple(
self.attention_layer_type
if (index + 1) % self.full_attention_interval == 0
else self.linear_layer_type
for index in range(self.num_hidden_layers)
)
@property
def rotary_dim(self) -> int:
return int(self.attention_head_dim * self.partial_rotary_factor)
@property
def effective_num_layers(self) -> int:
return self.num_hidden_layers + len(self.shared_layer_indices)
def to_dict(self) -> dict[str, Any]:
payload = asdict(self)
payload["layer_types"] = list(self.layer_types)
return payload
def save_json(self, path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(self.to_dict(), indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
@classmethod
def from_json(cls, path: Path) -> TinyGDNConfig:
payload = json.loads(path.read_text(encoding="utf-8"))
payload.pop("layer_types", None)
if "shared_layer_indices" in payload:
payload["shared_layer_indices"] = tuple(payload["shared_layer_indices"])
allowed = {item.name for item in fields(cls)}
return cls(**{key: value for key, value in payload.items() if key in allowed})
def haiku_config(**overrides: Any) -> TinyGDNConfig:
"""650M Haiku: 3×KDA + 1×Gated-MLA, AttnRes, SiTU-GLU, MTP."""
payload = {
"architecture": "HaikuForCausalLM",
"model_type": "haiku",
"vocab_size": 65_536,
"hidden_size": 1024,
"intermediate_size": 3840,
"num_hidden_layers": 36,
"num_attention_heads": 8,
"num_key_value_heads": 8,
"attention_head_dim": 128,
"full_attention_interval": 4,
"linear_num_heads": 8,
"linear_num_value_heads": 8,
"linear_head_dim": 128,
"linear_layer_type": "kda",
"attention_layer_type": "gated_mla",
"mlp_activation": "situ_glu",
"use_block_attn_res": True,
"mtp_num_heads": 2,
"mtp_adapter_rank": 128,
"mtp_loss_weight": 0.2,
"max_position_embeddings": 32_768,
"training_sequence_length": 2048,
}
payload.update(overrides)
return TinyGDNConfig(**payload)